257 lines
13 KiB
Go
257 lines
13 KiB
Go
package httpapi
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"testing"
|
|
|
|
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing"
|
|
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/identity"
|
|
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/jobs"
|
|
)
|
|
|
|
func TestBillingMemberRoutesUseRefreshedSessionScope(t *testing.T) {
|
|
service := &billingHTTPServiceStub{}
|
|
h := billingTestHandler(t, identity.Session{AuthMode: identity.AuthModeUser, User: identity.User{ID: "db-user", ClientID: "platform", OrganizationID: "db-org", OrganizationName: "DB Org", Role: "user"}}, service, &billingAccountStoreStub{})
|
|
|
|
response := serveJSON(t, h, http.MethodGet, "/api/billing", nil)
|
|
if response.Code != 200 || service.overviewOrganization != "db-org" || service.overviewAccount != "db-user" {
|
|
t.Fatalf("status=%d scope=%s/%s body=%s", response.Code, service.overviewOrganization, service.overviewAccount, response.Body.String())
|
|
}
|
|
response = serveJSON(t, h, http.MethodPost, "/api/billing/quote", map[string]any{"provider": "bailian", "capability": "image.generate", "ownerId": "attacker", "organizationId": "other"})
|
|
if response.Code != 200 || service.quote.AccountID != "db-user" || service.quote.OrganizationID != "db-org" || service.quote.Role != "user" {
|
|
t.Fatalf("status=%d quote=%+v body=%s", response.Code, service.quote, response.Body.String())
|
|
}
|
|
}
|
|
|
|
func TestBillingOverviewResponsesEncodeEmptyCollectionsAsArrays(t *testing.T) {
|
|
service := billing.NewService(&emptyBillingStore{}, nil)
|
|
member := billingTestHandler(t, identity.Session{AuthMode: identity.AuthModeUser, User: identity.User{ID: "user", ClientID: "platform", OrganizationID: "org", Role: "user"}}, service, nil)
|
|
admin := billingTestHandler(t, identity.Session{AuthMode: identity.AuthModeAdmin, User: identity.User{ID: "root", ClientID: "platform", Role: "super_admin"}}, service, nil)
|
|
|
|
assertJSONArrays(t, serveJSON(t, member, http.MethodGet, "/api/billing", nil), "ledger")
|
|
assertJSONArrays(t, serveJSON(t, admin, http.MethodGet, "/api/admin/billing", nil), "organizations", "members", "ledger", "priceRules")
|
|
}
|
|
|
|
func TestBillingQuotePreparesProviderAndModelFromServerOwnedEngineTargets(t *testing.T) {
|
|
service := &billingHTTPServiceStub{}
|
|
h := billingTestHandler(t, identity.Session{User: identity.User{ID: "db-user", ClientID: "platform", OrganizationID: "db-org", Role: "user"}}, service, nil)
|
|
h = NewBillingHandlerWithBuilder(h.(*billingHandler).authorizer, service, nil, jobs.ProviderJobBuilder{
|
|
ImageEngine: "bailian", ImageProvider: "bailian", ImageModel: "wan-image",
|
|
ImageEngines: map[string]jobs.ProviderTarget{"evolink": {Provider: "evolink", Model: "gpt-image-2"}},
|
|
VideoEngine: "seedance", VideoProvider: "seedance", VideoModel: "seedance-2",
|
|
NewID: func() string { return "quote-only" },
|
|
})
|
|
response := serveJSON(t, h, http.MethodPost, "/api/billing/quote", map[string]any{"engine": "evolink", "prompt": "hello", "quality": "high"})
|
|
if response.Code != 200 || service.quote.Provider != "evolink" || service.quote.Capability != "image.generate" || service.quote.ReqKey != "gpt-image-2" || service.quote.Parameters["quality"] != "high" {
|
|
t.Fatalf("status=%d quote=%+v body=%s", response.Code, service.quote, response.Body.String())
|
|
}
|
|
}
|
|
|
|
func TestBillingRequiresOrganizationAndMapsDomainStatuses(t *testing.T) {
|
|
service := &billingHTTPServiceStub{}
|
|
h := billingTestHandler(t, identity.Session{AuthMode: identity.AuthModeUser, User: identity.User{ID: "user", ClientID: "platform", Role: "user"}}, service, nil)
|
|
if got := serveJSON(t, h, http.MethodGet, "/api/billing", nil); got.Code != 422 {
|
|
t.Fatalf("status=%d body=%s", got.Code, got.Body.String())
|
|
}
|
|
|
|
service.err = &billing.StatusError{Status: 402, Err: billing.ErrInsufficientBalance}
|
|
h = billingTestHandler(t, identity.Session{AuthMode: identity.AuthModeUser, User: identity.User{ID: "user", ClientID: "platform", OrganizationID: "org", Role: "user"}}, service, nil)
|
|
if got := serveJSON(t, h, http.MethodPost, "/api/billing/quote", map[string]any{}); got.Code != 402 {
|
|
t.Fatalf("status=%d body=%s", got.Code, got.Body.String())
|
|
}
|
|
service.err = errors.New("database secret")
|
|
got := serveJSON(t, h, http.MethodGet, "/api/billing", nil)
|
|
if got.Code != 500 || bytes.Contains(got.Body.Bytes(), []byte("database secret")) {
|
|
t.Fatalf("status=%d body=%s", got.Code, got.Body.String())
|
|
}
|
|
}
|
|
|
|
func TestBillingAdminRoutesRequireSuperAdminAndValidateWrites(t *testing.T) {
|
|
service := &billingHTTPServiceStub{price: &billing.PriceRule{ID: "price-1", MarkupMultiplier: 1.2}}
|
|
account := &billingAccountStoreStub{}
|
|
orgAdmin := billingTestHandler(t, identity.Session{AuthMode: identity.AuthModeAdmin, User: identity.User{ID: "admin", ClientID: "platform", OrganizationID: "org", Role: "organization_admin"}}, service, account)
|
|
if got := serveJSON(t, orgAdmin, http.MethodGet, "/api/admin/billing", nil); got.Code != 403 {
|
|
t.Fatalf("organization admin status=%d", got.Code)
|
|
}
|
|
|
|
super := billingTestHandler(t, identity.Session{AuthMode: identity.AuthModeAdmin, User: identity.User{ID: "root", ClientID: "platform", Role: "super_admin"}}, service, account)
|
|
if got := serveJSON(t, super, http.MethodPatch, "/api/admin/billing/account", map[string]any{"accountName": " Acme ", "bankName": " Bank ", "accountNumber": " 123 ", "contact": " Ops "}); got.Code != 200 || account.saved.AccountName != "Acme" {
|
|
t.Fatalf("account status=%d saved=%+v body=%s", got.Code, account.saved, got.Body.String())
|
|
}
|
|
got := serveJSON(t, super, http.MethodPost, "/api/admin/billing/adjustments", map[string]any{"organizationId": "org", "amountYuan": 1.235, "direction": "debit", "note": " correction "})
|
|
if got.Code != 200 || service.adjustment.AmountFen != 124 || service.adjustment.DeltaFen != -124 || service.adjustment.OperatorID != "root" {
|
|
t.Fatalf("adjustment status=%d got=%+v body=%s", got.Code, service.adjustment, got.Body.String())
|
|
}
|
|
got = serveJSON(t, super, http.MethodPost, "/api/admin/billing/adjustments", map[string]any{"organizationId": "org", "amountYuan": "1.235", "direction": "credit", "note": " string amount "})
|
|
if got.Code != 200 || service.adjustment.AmountFen != 124 || service.adjustment.DeltaFen != 124 {
|
|
t.Fatalf("string adjustment status=%d got=%+v body=%s", got.Code, service.adjustment, got.Body.String())
|
|
}
|
|
var adjustmentResponse map[string]json.RawMessage
|
|
if err := json.Unmarshal(got.Body.Bytes(), &adjustmentResponse); err != nil {
|
|
t.Fatalf("decode adjustment response: %v", err)
|
|
}
|
|
if _, ok := adjustmentResponse["wallet"]; !ok {
|
|
t.Fatalf("adjustment response missing lowercase wallet: %s", got.Body.String())
|
|
}
|
|
if _, ok := adjustmentResponse["entry"]; !ok {
|
|
t.Fatalf("adjustment response missing lowercase entry: %s", got.Body.String())
|
|
}
|
|
if _, ok := adjustmentResponse["Wallet"]; ok {
|
|
t.Fatalf("adjustment response contains uppercase Wallet: %s", got.Body.String())
|
|
}
|
|
if _, ok := adjustmentResponse["Entry"]; ok {
|
|
t.Fatalf("adjustment response contains uppercase Entry: %s", got.Body.String())
|
|
}
|
|
if got := serveJSON(t, super, http.MethodPatch, "/api/admin/billing/prices/price-1", map[string]any{"standardUnitPriceFen": 1, "markupMultiplier": 2}); got.Code != 400 {
|
|
t.Fatalf("whitelist status=%d", got.Code)
|
|
}
|
|
if got := serveJSON(t, super, http.MethodPatch, "/api/admin/billing/prices/price-1", map[string]any{"markupMultiplier": .5}); got.Code != 400 {
|
|
t.Fatalf("multiplier status=%d", got.Code)
|
|
}
|
|
service.price.Dimensions = []billing.ParameterDimension{{Key: "quality", Tiers: []billing.ParameterTier{{Value: "high", Enabled: true}}}}
|
|
if got := serveJSON(t, super, http.MethodPatch, "/api/admin/billing/prices/price-1", map[string]any{"markupMultiplier": 2}); got.Code != 400 {
|
|
t.Fatalf("tier-required status=%d", got.Code)
|
|
}
|
|
if got := serveJSON(t, super, http.MethodPatch, "/api/admin/billing/prices/price-1", map[string]any{"markupMultiplier": 2, "dimensionKey": "quality", "tierValue": "missing"}); got.Code != 400 {
|
|
t.Fatalf("unknown tier status=%d", got.Code)
|
|
}
|
|
if got := serveJSON(t, super, http.MethodPatch, "/api/admin/billing/prices/price-1", map[string]any{"markupMultiplier": 2.123456, "dimensionKey": "quality", "tierValue": "high"}); got.Code != 200 || service.pricePatch.MarkupMultiplier != 2.1235 {
|
|
t.Fatalf("patch status=%d patch=%+v body=%s", got.Code, service.pricePatch, got.Body.String())
|
|
}
|
|
}
|
|
|
|
func TestBillingHandlerRejectsWrongMethods(t *testing.T) {
|
|
h := billingTestHandler(t, identity.Session{AuthMode: identity.AuthModeAdmin, User: identity.User{ID: "root", ClientID: "platform", Role: "super_admin"}}, &billingHTTPServiceStub{}, nil)
|
|
got := serveJSON(t, h, http.MethodDelete, "/api/billing", nil)
|
|
if got.Code != 405 || got.Header().Get("Allow") != http.MethodGet {
|
|
t.Fatalf("status=%d allow=%q", got.Code, got.Header().Get("Allow"))
|
|
}
|
|
}
|
|
|
|
type billingHTTPServiceStub struct {
|
|
overviewOrganization, overviewAccount string
|
|
quote billing.QuoteCommand
|
|
adjustment billing.AdjustmentCommand
|
|
pricePatch billing.PricePatch
|
|
price *billing.PriceRule
|
|
err error
|
|
}
|
|
|
|
func (s *billingHTTPServiceStub) Overview(_ context.Context, organizationID, accountID string) (billing.Overview, error) {
|
|
s.overviewOrganization, s.overviewAccount = organizationID, accountID
|
|
return billing.Overview{Wallet: billing.Wallet{OrganizationID: organizationID, Currency: billing.CurrencyCNY}}, s.err
|
|
}
|
|
func (s *billingHTTPServiceStub) Quote(_ context.Context, command billing.QuoteCommand) (*billing.Quote, error) {
|
|
s.quote = command
|
|
return &billing.Quote{Currency: billing.CurrencyCNY}, s.err
|
|
}
|
|
func (s *billingHTTPServiceStub) AdminOverview(context.Context) (billing.AdminOverview, error) {
|
|
return billing.AdminOverview{}, s.err
|
|
}
|
|
func (s *billingHTTPServiceStub) ListPrices(context.Context) ([]billing.PriceRule, error) {
|
|
return nil, s.err
|
|
}
|
|
func (s *billingHTTPServiceStub) GetPrice(_ context.Context, _ string) (*billing.PriceRule, error) {
|
|
return s.price, s.err
|
|
}
|
|
func (s *billingHTTPServiceStub) UpdatePrice(_ context.Context, _ string, patch billing.PricePatch) (*billing.PriceRule, error) {
|
|
s.pricePatch = patch
|
|
return s.price, s.err
|
|
}
|
|
func (s *billingHTTPServiceStub) Adjust(_ context.Context, command billing.AdjustmentCommand) (billing.AdjustmentResult, error) {
|
|
s.adjustment = command
|
|
return billing.AdjustmentResult{Wallet: billing.Wallet{OrganizationID: command.OrganizationID, Currency: billing.CurrencyCNY}}, s.err
|
|
}
|
|
|
|
type billingAccountStoreStub struct{ saved billing.AccountConfig }
|
|
|
|
func (s *billingAccountStoreStub) Load(context.Context) (billing.AccountConfig, error) {
|
|
return s.saved, nil
|
|
}
|
|
func (s *billingAccountStoreStub) Save(_ context.Context, value billing.AccountConfig) error {
|
|
s.saved = value
|
|
return nil
|
|
}
|
|
|
|
type emptyBillingStore struct{}
|
|
|
|
func (*emptyBillingStore) BillingWallet(context.Context, string) (billing.Wallet, error) {
|
|
return billing.Wallet{}, nil
|
|
}
|
|
func (*emptyBillingStore) BillingWallets(context.Context) ([]billing.Wallet, error) {
|
|
return nil, nil
|
|
}
|
|
func (*emptyBillingStore) BillingLedger(context.Context, string, string, int) ([]billing.LedgerEntry, error) {
|
|
return nil, nil
|
|
}
|
|
func (*emptyBillingStore) BillingOrganizations(context.Context) ([]billing.Organization, error) {
|
|
return nil, nil
|
|
}
|
|
func (*emptyBillingStore) BillingMembers(context.Context) ([]billing.Member, error) {
|
|
return nil, nil
|
|
}
|
|
func (*emptyBillingStore) BillingOrganizationExists(context.Context, string) (bool, error) {
|
|
return false, nil
|
|
}
|
|
func (*emptyBillingStore) ListBillingPriceRules(context.Context, bool) ([]billing.PriceRule, error) {
|
|
return nil, nil
|
|
}
|
|
func (*emptyBillingStore) GetBillingPriceRule(context.Context, string) (*billing.PriceRule, error) {
|
|
return nil, nil
|
|
}
|
|
func (*emptyBillingStore) UpdateBillingPriceRule(context.Context, string, billing.PricePatch) (*billing.PriceRule, error) {
|
|
return nil, nil
|
|
}
|
|
func (*emptyBillingStore) PostBillingWalletEntry(context.Context, billing.WalletPostParams) (billing.WalletPosting, error) {
|
|
return billing.WalletPosting{}, nil
|
|
}
|
|
|
|
func billingTestHandler(t *testing.T, session identity.Session, service BillingHTTPService, accounts BillingAccountStore) http.Handler {
|
|
t.Helper()
|
|
authorizer, err := NewPlatformAuthorizer(AuthState{Required: true, Configured: true}, &fixedSessionResolver{session: session})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return NewBillingHandler(authorizer, service, accounts)
|
|
}
|
|
|
|
type fixedSessionResolver struct{ session identity.Session }
|
|
|
|
func (r *fixedSessionResolver) Resolve(context.Context, string) (identity.Session, error) {
|
|
return r.session, nil
|
|
}
|
|
|
|
func serveJSON(t *testing.T, h http.Handler, method, path string, body any) *httptest.ResponseRecorder {
|
|
t.Helper()
|
|
var raw []byte
|
|
if body != nil {
|
|
raw, _ = json.Marshal(body)
|
|
}
|
|
req := httptest.NewRequest(method, path, bytes.NewReader(raw))
|
|
req.AddCookie(&http.Cookie{Name: identity.SessionCookieName, Value: "signed"})
|
|
response := httptest.NewRecorder()
|
|
h.ServeHTTP(response, req)
|
|
return response
|
|
}
|
|
|
|
func assertJSONArrays(t *testing.T, response *httptest.ResponseRecorder, fields ...string) {
|
|
t.Helper()
|
|
if response.Code != http.StatusOK {
|
|
t.Fatalf("status=%d body=%s", response.Code, response.Body.String())
|
|
}
|
|
var payload map[string]json.RawMessage
|
|
if err := json.Unmarshal(response.Body.Bytes(), &payload); err != nil {
|
|
t.Fatalf("decode response: %v", err)
|
|
}
|
|
for _, field := range fields {
|
|
if string(payload[field]) != "[]" {
|
|
t.Errorf("%s = %s, want [] (body=%s)", field, payload[field], response.Body.String())
|
|
}
|
|
}
|
|
}
|