diff --git a/backend/internal/billing/catalog.go b/backend/internal/billing/catalog.go new file mode 100644 index 0000000..157128d --- /dev/null +++ b/backend/internal/billing/catalog.go @@ -0,0 +1,253 @@ +package billing + +import ( + "errors" + "fmt" + "math" + "sort" + "strings" +) + +const CurrencyCNY = "CNY" + +var ( + ErrAmbiguousPriceRule = errors.New("billing price rules are ambiguous") + ErrPriceRuleNotFound = errors.New("billing price rule not found") + ErrParameterTier = errors.New("billing parameter tier not found") +) + +type Unit string + +const ( + UnitRequest Unit = "request" + UnitImage Unit = "image" + UnitVideoSecond Unit = "video_second" +) + +type QuantitySource string + +const ( + QuantityRequest QuantitySource = "request" + QuantityImageCount QuantitySource = "image_count" + QuantityDuration QuantitySource = "duration" +) + +type Parameters map[string]any +type Conditions map[string]any +type ConditionRange struct { + Min *float64 + Max *float64 + Values []any +} + +type ParameterTier struct { + Value any + Match any + StandardFactor, MarkupMultiplier float64 + Enabled bool +} +type ParameterDimension struct { + Key string + DefaultValue any + BaselineValue any + Tiers []ParameterTier +} +type PriceRule struct { + ID, Provider, Capability, ReqKey string + Unit Unit + QuantitySource QuantitySource + StandardUnitPriceFen int64 + MarkupMultiplier float64 + Enabled bool + Conditions Conditions + Priority int + Dimensions []ParameterDimension +} +type QuoteInput struct { + Provider, Capability, ReqKey, Source, Role string + Parameters Parameters + BillingDisabled bool +} +type Quote struct { + PriceRuleID string + Unit Unit + Quantity float64 + StandardUnitPriceFen, AmountFen int64 + MarkupMultiplier float64 + Currency string + QuotaExempt bool +} +type Catalog struct{ Rules []PriceRule } + +func (c Catalog) Quote(input QuoteInput) (*Quote, error) { + if input.BillingDisabled || input.Provider == "mock" { + return nil, nil + } + type candidate struct { + rule PriceRule + req, conditions, priority int + } + var matches []candidate + for _, rule := range c.Rules { + if !rule.Enabled || rule.Provider != input.Provider || rule.Capability != input.Capability || rule.ReqKey != "" && rule.ReqKey != input.ReqKey || !conditionsMatch(rule.Conditions, input.Parameters) { + continue + } + req := 0 + if rule.ReqKey != "" { + req = 1 + } + matches = append(matches, candidate{rule, req, len(rule.Conditions), rule.Priority}) + } + if len(matches) == 0 { + return nil, ErrPriceRuleNotFound + } + sort.Slice(matches, func(i, j int) bool { + a, b := matches[i], matches[j] + if a.req != b.req { + return a.req > b.req + } + if a.conditions != b.conditions { + return a.conditions > b.conditions + } + if a.priority != b.priority { + return a.priority > b.priority + } + return a.rule.ID < b.rule.ID + }) + w := matches[0] + if len(matches) > 1 && matches[1].req == w.req && matches[1].conditions == w.conditions && matches[1].priority == w.priority { + return nil, fmt.Errorf("%w: %s and %s", ErrAmbiguousPriceRule, w.rule.ID, matches[1].rule.ID) + } + price, markup, err := tierPrice(w.rule, input.Parameters) + if err != nil { + return nil, err + } + quantity := quantityFor(w.rule, input.Parameters) + return &Quote{PriceRuleID: w.rule.ID, Unit: w.rule.Unit, Quantity: quantity, StandardUnitPriceFen: price, MarkupMultiplier: markup, AmountFen: int64(math.Ceil(float64(price) * quantity * markup)), Currency: CurrencyCNY, QuotaExempt: input.Source == "platform" && input.Role == "super_admin"}, nil +} + +func tierPrice(rule PriceRule, params Parameters) (int64, float64, error) { + if len(rule.Dimensions) == 0 { + return rule.StandardUnitPriceFen, rule.MarkupMultiplier, nil + } + factor, markup := 1.0, 1.0 + for _, dimension := range rule.Dimensions { + actual, ok := params[dimension.Key] + if !ok { + actual = dimension.DefaultValue + if actual == nil { + actual = dimension.BaselineValue + } + } + found := false + for _, tier := range dimension.Tiers { + if tier.Enabled && conditionMatches(firstNonNil(tier.Match, tier.Value), actual) { + if tier.StandardFactor <= 0 || tier.MarkupMultiplier < 1 { + return 0, 0, ErrParameterTier + } + factor *= tier.StandardFactor + markup = math.Max(markup, tier.MarkupMultiplier) + found = true + break + } + } + if !found { + return 0, 0, ErrParameterTier + } + } + return int64(math.Ceil(float64(rule.StandardUnitPriceFen) * factor)), markup, nil +} +func quantityFor(rule PriceRule, params Parameters) float64 { + source := rule.QuantitySource + if source == "" { + if rule.Unit == UnitImage { + source = QuantityImageCount + } else if rule.Unit == UnitVideoSecond { + source = QuantityDuration + } else { + source = QuantityRequest + } + } + if source == QuantityRequest { + return 1 + } + key := "imageCount" + if source == QuantityDuration { + key = "duration" + } + n, ok := number(params[key]) + if !ok || n <= 0 { + return 1 + } + return math.Ceil(n) +} +func conditionsMatch(cs Conditions, p Parameters) bool { + for k, c := range cs { + if !conditionMatches(c, p[k]) { + return false + } + } + return true +} +func conditionMatches(condition, actual any) bool { + if actual == nil { + return false + } + switch c := condition.(type) { + case ConditionRange: + n, ok := number(actual) + if !ok { + return false + } + if c.Min != nil && n < *c.Min { + return false + } + if c.Max != nil && n > *c.Max { + return false + } + for _, v := range c.Values { + if scalarEqual(v, actual) { + return true + } + } + return len(c.Values) == 0 + case map[string]any: + n, ok := number(actual) + if min, yes := number(c["min"]); yes && (!ok || n < min) { + return false + } + if max, yes := number(c["max"]); yes && (!ok || n > max) { + return false + } + return true + default: + return scalarEqual(condition, actual) + } +} +func scalarEqual(a, b any) bool { + if x, ok := number(a); ok { + y, ok := number(b) + return ok && x == y + } + return strings.EqualFold(strings.TrimSpace(fmt.Sprint(a)), strings.TrimSpace(fmt.Sprint(b))) +} +func number(v any) (float64, bool) { + switch n := v.(type) { + case int: + return float64(n), true + case int64: + return float64(n), true + case float64: + return n, !math.IsNaN(n) && !math.IsInf(n, 0) + case float32: + return float64(n), true + default: + return 0, false + } +} +func firstNonNil(a, b any) any { + if a != nil { + return a + } + return b +} diff --git a/backend/internal/billing/catalog_test.go b/backend/internal/billing/catalog_test.go new file mode 100644 index 0000000..8b64c3c --- /dev/null +++ b/backend/internal/billing/catalog_test.go @@ -0,0 +1,78 @@ +package billing + +import ( + "encoding/json" + "errors" + "os" + "testing" +) + +func TestBillingCoreFixtureIsConsumedByGo(t *testing.T) { + data, err := os.ReadFile("../../../contracts/billing/core-v1.json") + if err != nil { + t.Fatal(err) + } + var fixture struct { + Currency string `json:"currency"` + MoneyUnit string `json:"moneyUnit"` + Statuses struct { + InsufficientBalance int `json:"insufficientBalance"` + IdempotencyConflict int `json:"idempotencyConflict"` + } `json:"statuses"` + IdempotencyKeys struct { + Charge string `json:"charge"` + Refund string `json:"refund"` + } `json:"idempotencyKeys"` + } + if err := json.Unmarshal(data, &fixture); err != nil { + t.Fatal(err) + } + if fixture.Currency != CurrencyCNY || fixture.MoneyUnit != "fen" || fixture.Statuses.InsufficientBalance != 402 || fixture.Statuses.IdempotencyConflict != 409 || fixture.IdempotencyKeys.Charge != "job-charge:{jobId}" || fixture.IdempotencyKeys.Refund != "job-refund:{jobId}" { + t.Fatalf("fixture = %#v", fixture) + } +} + +func TestCatalogQuotesMostSpecificRuleWithCeiledQuantityAndTierFactors(t *testing.T) { + catalog := Catalog{Rules: []PriceRule{ + {ID: "generic", Provider: "evolink", Capability: "image.generate", Unit: UnitImage, StandardUnitPriceFen: 10, MarkupMultiplier: 1.2, Enabled: true}, + {ID: "specific", Provider: "evolink", Capability: "image.generate", ReqKey: "gpt-image-2", Unit: UnitImage, QuantitySource: QuantityImageCount, StandardUnitPriceFen: 34, MarkupMultiplier: 1.2, Enabled: true, Conditions: Conditions{"quality": "high"}, Dimensions: []ParameterDimension{ + {Key: "quality", Tiers: []ParameterTier{{Value: "high", StandardFactor: 4, MarkupMultiplier: 1.5, Enabled: true}}}, + {Key: "resolution", DefaultValue: "1K", Tiers: []ParameterTier{{Value: "1K", StandardFactor: 2, MarkupMultiplier: 1.8, Enabled: true}}}, + }}, + }} + quote, err := catalog.Quote(QuoteInput{Provider: "evolink", Capability: "image.generate", ReqKey: "gpt-image-2", Parameters: Parameters{"quality": "HIGH", "imageCount": 1.2}}) + if err != nil { + t.Fatal(err) + } + if quote.PriceRuleID != "specific" || quote.Quantity != 2 || quote.StandardUnitPriceFen != 272 || quote.MarkupMultiplier != 1.8 || quote.AmountFen != 980 || quote.Currency != CurrencyCNY { + t.Fatalf("quote = %#v", quote) + } +} + +func TestCatalogRejectsAmbiguousWinners(t *testing.T) { + rules := []PriceRule{ + {ID: "a", Provider: "bailian", Capability: "image.generate", Unit: UnitRequest, StandardUnitPriceFen: 1, MarkupMultiplier: 1, Enabled: true}, + {ID: "b", Provider: "bailian", Capability: "image.generate", Unit: UnitRequest, StandardUnitPriceFen: 2, MarkupMultiplier: 1, Enabled: true}, + } + _, err := (Catalog{Rules: rules}).Quote(QuoteInput{Provider: "bailian", Capability: "image.generate"}) + if !errors.Is(err, ErrAmbiguousPriceRule) { + t.Fatalf("error = %v", err) + } +} + +func TestCatalogExemptsDisabledMockAndSuperAdmin(t *testing.T) { + rule := PriceRule{ID: "r", Provider: "bailian", Capability: "image.generate", Unit: UnitRequest, StandardUnitPriceFen: 10, MarkupMultiplier: 1, Enabled: true} + for _, input := range []QuoteInput{ + {BillingDisabled: true, Provider: "bailian", Capability: "image.generate"}, + {Provider: "mock", Capability: "image.generate"}, + } { + quote, err := (Catalog{Rules: []PriceRule{rule}}).Quote(input) + if err != nil || quote != nil { + t.Fatalf("Quote(%#v) = %#v, %v", input, quote, err) + } + } + quote, err := (Catalog{Rules: []PriceRule{rule}}).Quote(QuoteInput{Provider: "bailian", Capability: "image.generate", Source: "platform", Role: "super_admin"}) + if err != nil || quote == nil || !quote.QuotaExempt { + t.Fatalf("quote = %#v, %v", quote, err) + } +} diff --git a/backend/internal/billing/ledger.go b/backend/internal/billing/ledger.go new file mode 100644 index 0000000..f0fa1f6 --- /dev/null +++ b/backend/internal/billing/ledger.go @@ -0,0 +1,89 @@ +package billing + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "strings" + "time" +) + +var ( + ErrInsufficientBalance = errors.New("insufficient billing balance") + ErrIdempotencyConflict = errors.New("billing idempotency conflict") +) + +type WalletPostParams struct { + LedgerID, OrganizationID, AccountID, JobID, Kind string + DeltaFen int64 + Currency, IdempotencyKey, Description string + Metadata map[string]any +} +type WalletPosting struct { + LedgerID string + BalanceAfterFen, BalanceFen, TotalRechargedFen, TotalChargedFen, DeltaFen int64 + CreatedAt, UpdatedAt time.Time +} +type WalletPoster interface { + PostWalletEntry(context.Context, WalletPostParams) (WalletPosting, error) +} +type Ledger struct { + Poster WalletPoster + NewID func() string +} +type ChargeRequest struct { + OrganizationID, AccountID, JobID, Description string + AmountFen int64 + Metadata map[string]any +} +type RefundRequest = ChargeRequest + +type StatusError struct { + Status int + Err error +} + +func (e *StatusError) Error() string { return e.Err.Error() } +func (e *StatusError) Unwrap() error { return e.Err } +func HTTPStatus(err error) int { + var status *StatusError + if errors.As(err, &status) { + return status.Status + } + return 500 +} + +func (l Ledger) Charge(ctx context.Context, in ChargeRequest) (WalletPosting, error) { + return l.post(ctx, in, "charge", -in.AmountFen, "job-charge:"+in.JobID) +} +func (l Ledger) Refund(ctx context.Context, in RefundRequest) (WalletPosting, error) { + return l.post(ctx, in, "refund", in.AmountFen, "job-refund:"+in.JobID) +} +func (l Ledger) post(ctx context.Context, in ChargeRequest, kind string, delta int64, key string) (WalletPosting, error) { + if l.Poster == nil || l.NewID == nil { + return WalletPosting{}, errors.New("billing ledger is not configured") + } + if in.AmountFen <= 0 { + return WalletPosting{}, errors.New("billing amount must be positive") + } + posting, err := l.Poster.PostWalletEntry(ctx, WalletPostParams{LedgerID: l.NewID(), OrganizationID: in.OrganizationID, AccountID: in.AccountID, JobID: in.JobID, Kind: kind, DeltaFen: delta, Currency: CurrencyCNY, IdempotencyKey: key, Description: in.Description, Metadata: in.Metadata}) + if err == nil { + return posting, nil + } + message := err.Error() + if strings.Contains(message, "BILLING_INSUFFICIENT_BALANCE") { + return WalletPosting{}, &StatusError{402, fmt.Errorf("%w: %v", ErrInsufficientBalance, err)} + } + if strings.Contains(message, "BILLING_IDEMPOTENCY_PAYLOAD_MISMATCH") { + return WalletPosting{}, &StatusError{409, fmt.Errorf("%w: %v", ErrIdempotencyConflict, err)} + } + return WalletPosting{}, err +} + +func encodeMetadata(value map[string]any) (json.RawMessage, error) { + if value == nil { + return json.RawMessage(nil), nil + } + return json.Marshal(value) +} diff --git a/backend/internal/billing/ledger_test.go b/backend/internal/billing/ledger_test.go new file mode 100644 index 0000000..1ae25d9 --- /dev/null +++ b/backend/internal/billing/ledger_test.go @@ -0,0 +1,46 @@ +package billing + +import ( + "context" + "errors" + "testing" + "time" +) + +func TestLedgerChargesAndRefundsThroughWalletPosterWithFrozenKeys(t *testing.T) { + poster := &recordingPoster{result: WalletPosting{LedgerID: "ledger-1", CreatedAt: time.Unix(1, 0)}} + ledger := Ledger{Poster: poster, NewID: func() string { return "new-ledger" }} + charge, err := ledger.Charge(context.Background(), ChargeRequest{OrganizationID: "org-1", AccountID: "a", JobID: "job-1", AmountFen: 125}) + if err != nil || charge.LedgerID != "ledger-1" || poster.params.IdempotencyKey != "job-charge:job-1" || poster.params.DeltaFen != -125 || poster.params.Currency != "CNY" { + t.Fatalf("charge = %#v params=%#v err=%v", charge, poster.params, err) + } + _, err = ledger.Refund(context.Background(), RefundRequest{OrganizationID: "org-1", AccountID: "a", JobID: "job-1", AmountFen: 125}) + if err != nil || poster.params.IdempotencyKey != "job-refund:job-1" || poster.params.DeltaFen != 125 { + t.Fatalf("refund params=%#v err=%v", poster.params, err) + } +} + +func TestLedgerMapsWalletFailuresToFrozenStatuses(t *testing.T) { + for _, test := range []struct { + db error + want error + status int + }{{errors.New("BILLING_INSUFFICIENT_BALANCE"), ErrInsufficientBalance, 402}, {errors.New("BILLING_IDEMPOTENCY_PAYLOAD_MISMATCH"), ErrIdempotencyConflict, 409}} { + ledger := Ledger{Poster: &recordingPoster{err: test.db}, NewID: func() string { return "id" }} + _, err := ledger.Charge(context.Background(), ChargeRequest{OrganizationID: "o", JobID: "j", AmountFen: 1}) + if !errors.Is(err, test.want) || HTTPStatus(err) != test.status { + t.Fatalf("error=%v status=%d", err, HTTPStatus(err)) + } + } +} + +type recordingPoster struct { + params WalletPostParams + result WalletPosting + err error +} + +func (p *recordingPoster) PostWalletEntry(_ context.Context, v WalletPostParams) (WalletPosting, error) { + p.params = v + return p.result, p.err +} diff --git a/backend/internal/postgres/billing.go b/backend/internal/postgres/billing.go new file mode 100644 index 0000000..572f114 --- /dev/null +++ b/backend/internal/postgres/billing.go @@ -0,0 +1,68 @@ +package postgres + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + + "git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing" +) + +const ListBillingPriceRulesSQL = `SELECT id, provider, capability, req_key, unit, standard_unit_price_fen, markup_multiplier, enabled, conditions, quantity_source, priority, parameter_dimensions +FROM public.billing_price_rules +WHERE ($1::boolean OR enabled = true) +ORDER BY provider, capability, id` + +type BillingWalletPoster struct{ database *Database } + +func NewBillingWalletPoster(database *Database) BillingWalletPoster { + return BillingWalletPoster{database: database} +} +func (poster BillingWalletPoster) PostWalletEntry(ctx context.Context, p billing.WalletPostParams) (billing.WalletPosting, error) { + metadata, err := json.Marshal(p.Metadata) + if err != nil { + return billing.WalletPosting{}, fmt.Errorf("encode wallet metadata: %w", err) + } + row, err := poster.database.PostWalletEntry(ctx, WalletEntryParams{LedgerID: p.LedgerID, OrganizationID: p.OrganizationID, AccountID: p.AccountID, JobID: p.JobID, Kind: p.Kind, DeltaFen: p.DeltaFen, Currency: p.Currency, IdempotencyKey: p.IdempotencyKey, Description: p.Description, Metadata: metadata}) + if err != nil { + return billing.WalletPosting{}, err + } + return billing.WalletPosting{LedgerID: row.LedgerID, BalanceAfterFen: row.BalanceAfterFen, BalanceFen: row.BalanceFen, TotalRechargedFen: row.TotalRechargedFen, TotalChargedFen: row.TotalChargedFen, CreatedAt: row.CreatedAt, UpdatedAt: row.UpdatedAt, DeltaFen: row.DeltaFen}, nil +} + +func (db *Database) ListBillingPriceRules(ctx context.Context, includeDisabled bool) ([]billing.PriceRule, error) { + if db.config.Backend != BackendPostgres || db.querier == nil { + return nil, fmt.Errorf("PostgreSQL is unavailable when ZHINIAN_DATA_BACKEND=%s", db.config.Backend) + } + rows, err := db.querier.Query(ctx, ListBillingPriceRulesSQL, includeDisabled) + if err != nil { + return nil, fmt.Errorf("list billing price rules: %w", err) + } + defer rows.Close() + var out []billing.PriceRule + for rows.Next() { + var rule billing.PriceRule + var req, quantity sql.NullString + var conditions, dimensions json.RawMessage + if err := rows.Scan(&rule.ID, &rule.Provider, &rule.Capability, &req, &rule.Unit, &rule.StandardUnitPriceFen, &rule.MarkupMultiplier, &rule.Enabled, &conditions, &quantity, &rule.Priority, &dimensions); err != nil { + return nil, fmt.Errorf("scan billing price rule: %w", err) + } + rule.ReqKey = req.String + rule.QuantitySource = billing.QuantitySource(quantity.String) + if len(conditions) > 0 { + if err := json.Unmarshal(conditions, &rule.Conditions); err != nil { + return nil, fmt.Errorf("decode billing conditions: %w", err) + } + } + if len(dimensions) > 0 { + if err := json.Unmarshal(dimensions, &rule.Dimensions); err != nil { + return nil, fmt.Errorf("decode billing dimensions: %w", err) + } + } + out = append(out, rule) + } + return out, rows.Err() +} + +var _ billing.WalletPoster = BillingWalletPoster{} diff --git a/backend/internal/postgres/billing_test.go b/backend/internal/postgres/billing_test.go new file mode 100644 index 0000000..fbbb9ea --- /dev/null +++ b/backend/internal/postgres/billing_test.go @@ -0,0 +1,29 @@ +package postgres + +import ( + "context" + "reflect" + "testing" + + "git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing" +) + +func TestBillingWalletPosterDelegatesToExistingDatabaseFunction(t *testing.T) { + db := NewDatabase(Config{Backend: BackendPostgres}, &fakeQuerier{rows: [][]any{{"ledger-1", int64(75), int64(75), int64(100), int64(25), nil, nil, int64(-25)}}}) + got, err := NewBillingWalletPoster(db).PostWalletEntry(context.Background(), billing.WalletPostParams{LedgerID: "ledger-1", OrganizationID: "org", AccountID: "a", JobID: "job", Kind: "charge", DeltaFen: -25, Currency: "CNY", IdempotencyKey: "job-charge:job", Description: "charge", Metadata: map[string]any{"x": 1}}) + if err != nil || got.LedgerID != "ledger-1" || got.BalanceFen != 75 { + t.Fatalf("got=%#v err=%v", got, err) + } +} + +func TestListBillingPriceRulesUsesExplicitColumns(t *testing.T) { + q := &fakeQuerier{} + db := NewDatabase(Config{Backend: BackendPostgres}, q) + _, err := db.ListBillingPriceRules(context.Background(), false) + if err != nil { + t.Fatal(err) + } + if q.sql != ListBillingPriceRulesSQL || reflect.DeepEqual(q.args, []any{false}) == false { + t.Fatalf("sql=%q args=%#v", q.sql, q.args) + } +} diff --git a/backend/internal/postgres/usage.go b/backend/internal/postgres/usage.go new file mode 100644 index 0000000..33a2a1d --- /dev/null +++ b/backend/internal/postgres/usage.go @@ -0,0 +1,89 @@ +package postgres + +import ( + "context" + "database/sql" + "fmt" + + "git.nianxx.cn/wangxuming/NianAIGC/backend/internal/usage" +) + +const InsertUsageEventSQL = `INSERT INTO public.usage_events (id, owner_id, job_id, source, capability, provider, req_key, account_username, account_display_name, tenant_id, organization_id, organization_name, quantity, estimated_unit, charged_amount_fen, currency, created_at) +VALUES ($1::text, $2::text, $3::text, $4::text, $5::text, NULLIF($6::text, ''), NULLIF($7::text, ''), NULLIF($8::text, ''), NULLIF($9::text, ''), NULLIF($10::text, ''), NULLIF($11::text, ''), NULLIF($12::text, ''), $13::integer, $14::text, $15::bigint, NULLIF($16::text, ''), $17::timestamptz) +ON CONFLICT (job_id) DO NOTHING +RETURNING id, owner_id, job_id, source, capability, COALESCE(provider, ''), COALESCE(req_key, ''), COALESCE(account_username, ''), COALESCE(account_display_name, ''), COALESCE(tenant_id, ''), COALESCE(organization_id, ''), COALESCE(organization_name, ''), quantity, estimated_unit, charged_amount_fen, COALESCE(currency, ''), created_at::text` +const ListUsageEventsSQL = `SELECT id, owner_id, job_id, source, capability, COALESCE(provider, ''), COALESCE(req_key, ''), COALESCE(account_username, ''), COALESCE(account_display_name, ''), COALESCE(tenant_id, ''), COALESCE(organization_id, ''), COALESCE(organization_name, ''), quantity, estimated_unit, charged_amount_fen, COALESCE(currency, ''), created_at::text +FROM public.usage_events +WHERE ($1::text = '' OR owner_id = $1::text) AND ($2::text = '' OR organization_id = $2::text) AND ($3::timestamptz IS NULL OR created_at >= $3::timestamptz) AND ($4::timestamptz IS NULL OR created_at < $4::timestamptz) +ORDER BY created_at DESC, id DESC` + +func (db *Database) InsertUsageEvent(ctx context.Context, e usage.Event) (usage.Event, bool, error) { + if db.config.Backend != BackendPostgres || db.querier == nil { + return usage.Event{}, false, fmt.Errorf("PostgreSQL is unavailable when ZHINIAN_DATA_BACKEND=%s", db.config.Backend) + } + rows, err := db.querier.Query(ctx, InsertUsageEventSQL, e.ID, e.OwnerID, e.JobID, e.Source, e.Capability, e.Provider, e.ReqKey, e.AccountUsername, e.AccountDisplayName, e.TenantID, e.OrganizationID, e.OrganizationName, e.Quantity, e.EstimatedUnit, e.ChargedAmountFen, e.Currency, e.CreatedAt) + if err != nil { + return usage.Event{}, false, fmt.Errorf("insert usage event: %w", err) + } + defer rows.Close() + if !rows.Next() { + return usage.Event{}, false, rows.Err() + } + out, err := scanUsage(rows) + return out, err == nil, err +} +func (db *Database) ListUsageEvents(ctx context.Context, f usage.Filters) ([]usage.Event, error) { + if db.config.Backend != BackendPostgres || db.querier == nil { + return nil, fmt.Errorf("PostgreSQL is unavailable when ZHINIAN_DATA_BACKEND=%s", db.config.Backend) + } + rows, err := db.querier.Query(ctx, ListUsageEventsSQL, f.OwnerID, f.OrganizationID, nullableTime(f.From), nullableTime(f.To)) + if err != nil { + return nil, fmt.Errorf("list usage events: %w", err) + } + defer rows.Close() + var out []usage.Event + for rows.Next() { + event, err := scanUsage(rows) + if err != nil { + return nil, err + } + if f.Capability != "" && event.Capability != f.Capability || f.Provider != "" && event.Provider != f.Provider { + continue + } + out = append(out, event) + } + return out, rows.Err() +} +func scanUsage(rows Rows) (usage.Event, error) { + var e usage.Event + var charged sql.NullInt64 + if err := rows.Scan(&e.ID, &e.OwnerID, &e.JobID, &e.Source, &e.Capability, &e.Provider, &e.ReqKey, &e.AccountUsername, &e.AccountDisplayName, &e.TenantID, &e.OrganizationID, &e.OrganizationName, &e.Quantity, &e.EstimatedUnit, &charged, &e.Currency, &e.CreatedAt); err != nil { + return usage.Event{}, fmt.Errorf("scan usage event: %w", err) + } + if charged.Valid { + e.ChargedAmountFen = &charged.Int64 + } + return e, nil +} +func nullableTime(v string) any { + if v == "" { + return nil + } + return v +} + +type UsageRepository struct{ database *Database } + +func NewUsageRepository(database *Database) UsageRepository { + return UsageRepository{database: database} +} + +func (repository UsageRepository) Insert(event usage.Event) (usage.Event, bool, error) { + return repository.database.InsertUsageEvent(context.Background(), event) +} + +func (repository UsageRepository) List(filters usage.Filters) ([]usage.Event, error) { + return repository.database.ListUsageEvents(context.Background(), filters) +} + +var _ usage.Repository = UsageRepository{} diff --git a/backend/internal/postgres/usage_test.go b/backend/internal/postgres/usage_test.go new file mode 100644 index 0000000..5ef07cd --- /dev/null +++ b/backend/internal/postgres/usage_test.go @@ -0,0 +1,33 @@ +package postgres + +import ( + "context" + "strings" + "testing" + + "git.nianxx.cn/wangxuming/NianAIGC/backend/internal/usage" +) + +func TestInsertUsageEventUsesJobConflictForDedupe(t *testing.T) { + q := &fakeQuerier{} + db := NewDatabase(Config{Backend: BackendPostgres}, q) + _, inserted, err := db.InsertUsageEvent(context.Background(), usage.Event{ID: "e", JobID: "j", OwnerID: "a", Source: "platform", Quantity: 1}) + if err != nil || inserted { + t.Fatalf("inserted=%v err=%v", inserted, err) + } + if q.sql != InsertUsageEventSQL || !strings.Contains(q.sql, "ON CONFLICT (job_id) DO NOTHING") { + t.Fatalf("sql=%q", q.sql) + } +} + +func TestListUsageEventsUsesExplicitColumnsAndFilters(t *testing.T) { + q := &fakeQuerier{} + db := NewDatabase(Config{Backend: BackendPostgres}, q) + _, err := db.ListUsageEvents(context.Background(), usage.Filters{OwnerID: "a", OrganizationID: "o", From: "2026-01-01", To: "2026-02-01"}) + if err != nil { + t.Fatal(err) + } + if q.sql != ListUsageEventsSQL || len(q.args) != 4 { + t.Fatalf("sql=%q args=%#v", q.sql, q.args) + } +} diff --git a/backend/internal/usage/usage.go b/backend/internal/usage/usage.go new file mode 100644 index 0000000..d18cded --- /dev/null +++ b/backend/internal/usage/usage.go @@ -0,0 +1,99 @@ +package usage + +import ( + "errors" + "fmt" + "sort" + "time" +) + +type Preset string + +const ( + PresetToday Preset = "today" + Preset7Days Preset = "7d" + Preset30Days Preset = "30d" + PresetMonth Preset = "month" +) + +var ErrForbidden = errors.New("usage scope forbidden") +var shanghai = time.FixedZone("Asia/Shanghai", 8*60*60) + +type DateRange struct { + From, To, StartDate, EndDate string + DayCount int +} + +func PresetRange(preset Preset, now time.Time) (DateRange, error) { + today := now.In(shanghai) + start := time.Date(today.Year(), today.Month(), today.Day(), 0, 0, 0, 0, shanghai) + switch preset { + case PresetToday: + case Preset7Days: + start = start.AddDate(0, 0, -6) + case Preset30Days: + start = start.AddDate(0, 0, -29) + case PresetMonth: + start = time.Date(today.Year(), today.Month(), 1, 0, 0, 0, 0, shanghai) + default: + return DateRange{}, fmt.Errorf("unknown usage preset %q", preset) + } + end := time.Date(today.Year(), today.Month(), today.Day(), 0, 0, 0, 0, shanghai) + return DateRange{From: start.Format(time.RFC3339), To: end.AddDate(0, 0, 1).Format(time.RFC3339), StartDate: start.Format("2006-01-02"), EndDate: end.Format("2006-01-02"), DayCount: int(end.Sub(start).Hours()/24) + 1}, nil +} + +type Event struct { + ID, OwnerID, JobID, Source, Capability, Provider, ReqKey, AccountUsername, AccountDisplayName, TenantID, OrganizationID, OrganizationName, EstimatedUnit, Currency, CreatedAt string + Quantity int + ChargedAmountFen *int64 +} +type Filters struct{ OwnerID, OrganizationID, Capability, Provider, From, To string } +type Requester struct{ AccountID, OrganizationID, Role string } +type Repository interface { + Insert(Event) (Event, bool, error) + List(Filters) ([]Event, error) +} +type Service struct{ Repository Repository } + +func (s Service) Record(event Event) (*Event, error) { + if event.Source == "api" || event.Provider == "mock" { + return nil, nil + } + saved, _, err := s.Repository.Insert(event) + if err != nil { + return nil, err + } + return &saved, nil +} +func (s Service) Report(who Requester, filters Filters) ([]Event, error) { + switch who.Role { + case "super_admin": + case "organization_admin": + if filters.OrganizationID != "" && filters.OrganizationID != who.OrganizationID { + return nil, ErrForbidden + } + filters.OrganizationID = who.OrganizationID + case "user": + filters.OwnerID = who.AccountID + default: + filters.OwnerID = who.AccountID + } + events, err := s.Repository.List(filters) + if err != nil { + return nil, err + } + seen := map[string]bool{} + out := make([]Event, 0, len(events)) + for _, e := range events { + if e.Source == "api" || e.Provider == "mock" || seen[e.JobID] || filters.OwnerID != "" && e.OwnerID != filters.OwnerID || filters.OrganizationID != "" && e.OrganizationID != filters.OrganizationID || filters.Capability != "" && e.Capability != filters.Capability || filters.Provider != "" && e.Provider != filters.Provider { + continue + } + seen[e.JobID] = true + if who.Role == "organization_admin" { + e.AccountUsername = "" + } + out = append(out, e) + } + sort.SliceStable(out, func(i, j int) bool { return out[i].CreatedAt > out[j].CreatedAt }) + return out, nil +} diff --git a/backend/internal/usage/usage_test.go b/backend/internal/usage/usage_test.go new file mode 100644 index 0000000..99ab0ff --- /dev/null +++ b/backend/internal/usage/usage_test.go @@ -0,0 +1,91 @@ +package usage + +import ( + "encoding/json" + "errors" + "os" + "testing" + "time" +) + +func TestUsageCoreFixtureIsConsumedByGo(t *testing.T) { + data, err := os.ReadFile("../../../contracts/usage/core-v1.json") + if err != nil { + t.Fatal(err) + } + var fixture struct { + TimeZone string `json:"timeZone"` + DedupeKey string `json:"dedupeKey"` + Excluded struct { + Sources []string `json:"sources"` + Providers []string `json:"providers"` + } `json:"excluded"` + } + if err := json.Unmarshal(data, &fixture); err != nil { + t.Fatal(err) + } + if fixture.TimeZone != "Asia/Shanghai" || fixture.DedupeKey != "jobId" || len(fixture.Excluded.Sources) != 1 || fixture.Excluded.Sources[0] != "api" || len(fixture.Excluded.Providers) != 1 || fixture.Excluded.Providers[0] != "mock" { + t.Fatalf("fixture = %#v", fixture) + } +} + +func TestPresetUsesShanghaiCalendarBoundaries(t *testing.T) { + rangeValue, err := PresetRange(PresetMonth, time.Date(2026, 7, 27, 16, 30, 0, 0, time.UTC)) + if err != nil { + t.Fatal(err) + } + if rangeValue.StartDate != "2026-07-01" || rangeValue.EndDate != "2026-07-28" || rangeValue.From != "2026-07-01T00:00:00+08:00" || rangeValue.To != "2026-07-29T00:00:00+08:00" || rangeValue.DayCount != 28 { + t.Fatalf("range = %#v", rangeValue) + } +} + +func TestServiceDedupesEligibleEventsAndScopesReports(t *testing.T) { + repo := &memoryRepository{} + service := Service{Repository: repo} + for _, event := range []Event{ + {ID: "1", JobID: "job-1", OwnerID: "a", Source: "platform", Provider: "bailian", OrganizationID: "org-1", AccountUsername: "secret", CreatedAt: "2026-07-03T01:00:00Z"}, + {ID: "2", JobID: "job-1", OwnerID: "a", Source: "platform", Provider: "bailian", OrganizationID: "org-1", CreatedAt: "2026-07-03T01:01:00Z"}, + {ID: "3", JobID: "job-api", OwnerID: "a", Source: "api", Provider: "bailian", OrganizationID: "org-1", CreatedAt: "2026-07-03T01:02:00Z"}, + {ID: "4", JobID: "job-mock", OwnerID: "a", Source: "platform", Provider: "mock", OrganizationID: "org-1", CreatedAt: "2026-07-03T01:03:00Z"}, + {ID: "5", JobID: "job-2", OwnerID: "b", Source: "platform", Provider: "seedance", OrganizationID: "org-2", AccountUsername: "other-secret", CreatedAt: "2026-07-03T01:04:00Z"}, + } { + if _, err := service.Record(event); err != nil { + t.Fatal(err) + } + } + + personal, err := service.Report(Requester{AccountID: "a", Role: "user"}, Filters{}) + if err != nil || len(personal) != 1 || personal[0].OwnerID != "a" { + t.Fatalf("personal = %#v, %v", personal, err) + } + admin, err := service.Report(Requester{AccountID: "admin", OrganizationID: "org-1", Role: "organization_admin"}, Filters{}) + if err != nil || len(admin) != 1 || admin[0].OrganizationID != "org-1" || admin[0].AccountUsername != "" { + t.Fatalf("admin = %#v, %v", admin, err) + } + super, err := service.Report(Requester{AccountID: "root", Role: "super_admin"}, Filters{}) + if err != nil || len(super) != 2 || super[0].AccountUsername == "" { + t.Fatalf("super = %#v, %v", super, err) + } +} + +func TestOrganizationAdminCannotEscapeOrganizationScope(t *testing.T) { + _, err := (Service{Repository: &memoryRepository{}}).Report(Requester{Role: "organization_admin", OrganizationID: "org-1"}, Filters{OrganizationID: "org-2"}) + if !errors.Is(err, ErrForbidden) { + t.Fatalf("error = %v", err) + } +} + +type memoryRepository struct{ events []Event } + +func (r *memoryRepository) Insert(event Event) (Event, bool, error) { + for _, current := range r.events { + if current.JobID == event.JobID { + return current, false, nil + } + } + r.events = append(r.events, event) + return event, true, nil +} +func (r *memoryRepository) List(filters Filters) ([]Event, error) { + return append([]Event(nil), r.events...), nil +} diff --git a/contracts/billing/core-v1.json b/contracts/billing/core-v1.json new file mode 100644 index 0000000..b480a86 --- /dev/null +++ b/contracts/billing/core-v1.json @@ -0,0 +1,15 @@ +{ + "version": 1, + "currency": "CNY", + "moneyUnit": "fen", + "statuses": { "insufficientBalance": 402, "idempotencyConflict": 409 }, + "idempotencyKeys": { "charge": "job-charge:{jobId}", "refund": "job-refund:{jobId}" }, + "quote": { + "quantityRounding": "ceil", + "ruleOrder": ["reqKeySpecificity", "conditionSpecificity", "priority"], + "tiedWinner": "ambiguity_error", + "tierStandardFactors": "multiply", + "tierMarkupMultipliers": "maximum" + }, + "exemptions": { "billingDisabled": true, "mockProvider": true, "platformSuperAdminQuota": true } +} diff --git a/contracts/usage/core-v1.json b/contracts/usage/core-v1.json new file mode 100644 index 0000000..bfe8859 --- /dev/null +++ b/contracts/usage/core-v1.json @@ -0,0 +1,12 @@ +{ + "version": 1, + "timeZone": "Asia/Shanghai", + "excluded": { "sources": ["api"], "providers": ["mock"] }, + "dedupeKey": "jobId", + "scopes": { + "personal": "accountId", + "organization_admin": "organizationId", + "super_admin": "all" + }, + "organizationAdminRedaction": ["accountUsername"] +} diff --git a/tests/billing-go-core-contract.test.ts b/tests/billing-go-core-contract.test.ts new file mode 100644 index 0000000..7d002ba --- /dev/null +++ b/tests/billing-go-core-contract.test.ts @@ -0,0 +1,15 @@ +import { readFile } from "node:fs/promises"; +import { describe, expect, it } from "vitest"; + +describe("Go billing core compatibility fixture", () => { + it("freezes money, errors, idempotency, quote ordering, and exemptions", async () => { + const fixture = JSON.parse(await readFile("contracts/billing/core-v1.json", "utf8")); + expect(fixture).toMatchObject({ + currency: "CNY", moneyUnit: "fen", + statuses: { insufficientBalance: 402, idempotencyConflict: 409 }, + idempotencyKeys: { charge: "job-charge:{jobId}", refund: "job-refund:{jobId}" }, + quote: { quantityRounding: "ceil", tiedWinner: "ambiguity_error", tierStandardFactors: "multiply", tierMarkupMultipliers: "maximum" }, + exemptions: { billingDisabled: true, mockProvider: true, platformSuperAdminQuota: true } + }); + }); +}); diff --git a/tests/usage-go-core-contract.test.ts b/tests/usage-go-core-contract.test.ts new file mode 100644 index 0000000..bf1cdce --- /dev/null +++ b/tests/usage-go-core-contract.test.ts @@ -0,0 +1,14 @@ +import { readFile } from "node:fs/promises"; +import { describe, expect, it } from "vitest"; + +describe("Go usage core compatibility fixture", () => { + it("freezes Shanghai boundaries, exclusions, dedupe, scope, and redaction", async () => { + const fixture = JSON.parse(await readFile("contracts/usage/core-v1.json", "utf8")); + expect(fixture).toEqual({ + version: 1, timeZone: "Asia/Shanghai", + excluded: { sources: ["api"], providers: ["mock"] }, dedupeKey: "jobId", + scopes: { personal: "accountId", organization_admin: "organizationId", super_admin: "all" }, + organizationAdminRedaction: ["accountUsername"] + }); + }); +});