feat: add billing and usage domain cores
This commit is contained in:
1 parent
48dd5d07c8
commit
5e60bb40e7
14 files changed
+931
No files matched your search
@@ -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{}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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{}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user