Files
cmhub/apps/billing/tests.py
T

710 lines
26 KiB
Python

import threading
import time
import uuid
from datetime import timedelta
from decimal import Decimal
from cryptography.fernet import Fernet
from django.contrib import admin
from django.contrib.auth import get_user_model
from django.core.exceptions import ValidationError
from django.db import IntegrityError, OperationalError, connections, transaction
from django.test import TestCase, TransactionTestCase, override_settings
from django.utils import timezone
from apps.ai.models import AiModel, ModelAlias
from apps.billing.models import (
CallRecord,
ExchangeRate,
PointsLedger,
PricingRule,
RechargeOrder,
)
from apps.billing.pricing import (
NoPricingRuleError,
calculate_points_cost,
calculate_points_granted,
get_pricing_rule,
quote_recharge_points,
)
from apps.billing.services import (
InsufficientPointsError,
InvalidCallStateError,
RechargeAmountMismatchError,
RechargePayment,
apply_recharge_payment,
mark_call_success,
precharge_call,
query_and_apply_recharge_payment,
refund_call_points,
)
from apps.users.models import ApiKey, UserWallet
TEST_ENCRYPTION_KEY = Fernet.generate_key().decode("ascii")
class BillingCoreModelTests(TestCase):
def setUp(self):
suffix = uuid.uuid4().hex[:8]
self.user = get_user_model().objects.create_user(
username=f"client-{suffix}",
email=f"client-{suffix}@example.com",
password="password",
)
def test_user_wallet_defaults_to_zero_and_rejects_negative_balance(self):
wallet = UserWallet.objects.create(user=self.user)
self.assertEqual(wallet.points_balance, 0)
other_user = get_user_model().objects.create_user(
username=f"negative-{uuid.uuid4().hex[:8]}",
email=f"negative-{uuid.uuid4().hex[:8]}@example.com",
password="password",
)
with self.assertRaises(IntegrityError):
with transaction.atomic():
UserWallet.objects.create(user=other_user, points_balance=-1)
def test_api_key_hashes_plaintext_and_matches_only_raw_key(self):
api_key, raw_key = ApiKey.create_for_user(self.user, name="desktop")
self.assertTrue(raw_key.startswith("sk_cmhub_"))
self.assertEqual(api_key.key_prefix, raw_key[: ApiKey.KEY_PREFIX_LENGTH])
self.assertEqual(len(api_key.key_hash), 64)
self.assertNotEqual(api_key.key_hash, raw_key)
self.assertNotIn(raw_key, str(api_key.__dict__))
self.assertTrue(api_key.matches_key(raw_key))
self.assertFalse(api_key.matches_key(raw_key + "-wrong"))
self.assertEqual(api_key.status, ApiKey.Status.ACTIVE)
def test_call_record_stores_summary_reference_but_no_provider_raw_field(self):
api_key, _raw_key = ApiKey.create_for_user(self.user)
call = CallRecord.objects.create(
user=self.user,
api_key=api_key,
operation_type=CallRecord.OperationType.IMAGE,
alias="image-hd",
model_used="gpt-image-2",
resolution="1K",
prompt="Generate an image",
points_cost=10,
status=CallRecord.Status.SUCCESS,
upstream_latency_ms=1234,
result_ref="https://cdn.example.test/result.png",
result_summary="stored image result",
)
field_names = {field.name for field in CallRecord._meta.fields}
self.assertNotIn("raw", field_names)
self.assertNotIn("provider_raw", field_names)
self.assertEqual(call.user, self.user)
self.assertEqual(call.api_key, api_key)
self.assertEqual(call.alias, "image-hd")
self.assertEqual(call.model_used, "gpt-image-2")
self.assertEqual(call.result_ref, "https://cdn.example.test/result.png")
self.assertEqual(call.result_summary, "stored image result")
def test_points_ledger_records_balance_and_requires_adjust_reason(self):
call = CallRecord.objects.create(
user=self.user,
operation_type=CallRecord.OperationType.TITLE,
alias="title-standard",
model_used="gpt-5.5",
points_cost=2,
status=CallRecord.Status.SUCCESS,
)
ledger = PointsLedger.objects.create(
user=self.user,
change_type=PointsLedger.ChangeType.CONSUME,
points_delta=-2,
balance_after=98,
ref_call=call,
)
self.assertEqual(ledger.ref_call, call)
self.assertEqual(ledger.balance_after, 98)
adjustment = PointsLedger(
user=self.user,
change_type=PointsLedger.ChangeType.ADJUST,
points_delta=10,
balance_after=108,
)
with self.assertRaises(ValidationError):
adjustment.full_clean()
def test_points_ledger_rejects_zero_delta_and_negative_balance_after(self):
with self.assertRaises(IntegrityError):
with transaction.atomic():
PointsLedger.objects.create(
user=self.user,
change_type=PointsLedger.ChangeType.RECHARGE,
points_delta=0,
balance_after=100,
)
with self.assertRaises(IntegrityError):
with transaction.atomic():
PointsLedger.objects.create(
user=self.user,
change_type=PointsLedger.ChangeType.CONSUME,
points_delta=-1,
balance_after=-1,
)
def test_points_ledger_allows_consume_and_refund_but_rejects_duplicate_refund(self):
call = CallRecord.objects.create(
user=self.user,
operation_type=CallRecord.OperationType.TITLE,
alias="title-standard",
model_used="gpt-5.5",
points_cost=2,
status=CallRecord.Status.FAILED,
)
PointsLedger.objects.create(
user=self.user,
change_type=PointsLedger.ChangeType.CONSUME,
points_delta=-2,
balance_after=98,
ref_call=call,
)
PointsLedger.objects.create(
user=self.user,
change_type=PointsLedger.ChangeType.REFUND,
points_delta=2,
balance_after=100,
ref_call=call,
)
with self.assertRaises(IntegrityError):
with transaction.atomic():
PointsLedger.objects.create(
user=self.user,
change_type=PointsLedger.ChangeType.REFUND,
points_delta=2,
balance_after=102,
ref_call=call,
)
def test_points_ledger_rejects_duplicate_recharge_for_same_order(self):
order = RechargeOrder.objects.create(
user=self.user,
order_no=f"R{uuid.uuid4().hex[:12]}",
amount_money=Decimal("20.00"),
pay_method=RechargeOrder.PayMethod.WEIXIN,
exchange_rate=Decimal("10.0000"),
points_granted=200,
)
other_order = RechargeOrder.objects.create(
user=self.user,
order_no=f"R{uuid.uuid4().hex[:12]}",
amount_money=Decimal("30.00"),
pay_method=RechargeOrder.PayMethod.WEIXIN,
exchange_rate=Decimal("10.0000"),
points_granted=300,
)
PointsLedger.objects.create(
user=self.user,
change_type=PointsLedger.ChangeType.RECHARGE,
points_delta=200,
balance_after=200,
ref_order_id=order.id,
)
PointsLedger.objects.create(
user=self.user,
change_type=PointsLedger.ChangeType.RECHARGE,
points_delta=300,
balance_after=500,
ref_order_id=other_order.id,
)
with self.assertRaises(IntegrityError):
with transaction.atomic():
PointsLedger.objects.create(
user=self.user,
change_type=PointsLedger.ChangeType.RECHARGE,
points_delta=200,
balance_after=700,
ref_order_id=order.id,
)
def test_billing_models_are_registered_in_admin(self):
self.assertIn(UserWallet, admin.site._registry)
self.assertIn(ApiKey, admin.site._registry)
self.assertIn(PointsLedger, admin.site._registry)
self.assertIn(CallRecord, admin.site._registry)
self.assertIn(RechargeOrder, admin.site._registry)
@override_settings(AI_KEY_ENCRYPTION_KEY=TEST_ENCRYPTION_KEY)
class PricingCalculationTests(TestCase):
def setUp(self):
suffix = uuid.uuid4().hex[:8]
self.title_alias = f"title-standard-{suffix}"
self.image_alias = f"image-hd-{suffix}"
self.text_model = self.create_ai_model(
name="Text model A",
model=f"gpt-5.5-a-{suffix}",
capabilities=["text"],
)
self.replacement_text_model = self.create_ai_model(
name="Text model B",
model=f"gpt-5.5-b-{suffix}",
capabilities=["text"],
)
self.alias = ModelAlias.objects.create(
operation_type=ModelAlias.OperationType.TITLE,
alias=self.title_alias,
ai_model=self.text_model,
)
def create_ai_model(self, *, name, model, capabilities):
ai_model = AiModel(
name=name,
url="https://api.vectorengine.ai/v1",
model=model,
api_type=AiModel.ApiType.CHAT,
capabilities=capabilities,
)
ai_model.set_api_key("sk-test-secret")
ai_model.save()
return ai_model
def test_exact_resolution_rule_overrides_alias_default(self):
PricingRule.objects.create(
operation_type=CallRecord.OperationType.IMAGE,
alias=self.image_alias,
resolution="",
points_cost=10,
)
PricingRule.objects.create(
operation_type=CallRecord.OperationType.IMAGE,
alias=self.image_alias,
resolution="1k",
points_cost=15,
)
self.assertEqual(
calculate_points_cost(
CallRecord.OperationType.IMAGE,
self.image_alias,
"1K",
),
15,
)
self.assertEqual(
calculate_points_cost(
CallRecord.OperationType.IMAGE,
self.image_alias,
"2K",
),
10,
)
exact_rule = get_pricing_rule(
CallRecord.OperationType.IMAGE,
self.image_alias,
"1k",
)
self.assertEqual(exact_rule.resolution, "1K")
def test_pricing_is_bound_to_alias_not_underlying_model(self):
PricingRule.objects.create(
operation_type=CallRecord.OperationType.TITLE,
alias=self.title_alias,
points_cost=2,
)
self.assertEqual(
calculate_points_cost(
CallRecord.OperationType.TITLE,
self.title_alias,
"1K",
),
2,
)
self.alias.ai_model = self.replacement_text_model
self.alias.save()
self.assertEqual(
calculate_points_cost(
CallRecord.OperationType.TITLE,
self.title_alias,
"1K",
),
2,
)
def test_missing_pricing_rule_raises_no_pricing_rule_code(self):
with self.assertRaises(NoPricingRuleError) as context:
calculate_points_cost(
CallRecord.OperationType.TITLE,
"missing-alias",
"1K",
)
self.assertEqual(context.exception.code, "no_pricing_rule")
def test_exchange_rate_uses_latest_active_effective_rate_and_floors_points(self):
now = timezone.now()
ExchangeRate.objects.create(
currency="cny",
points_per_unit=Decimal("10.0000"),
effective_from=now - timedelta(days=2),
)
current = ExchangeRate.objects.create(
currency="CNY",
points_per_unit=Decimal("12.5000"),
effective_from=now - timedelta(days=1),
)
ExchangeRate.objects.create(
currency="CNY",
points_per_unit=Decimal("20.0000"),
effective_from=now,
is_active=False,
)
ExchangeRate.objects.create(
currency="CNY",
points_per_unit=Decimal("99.0000"),
effective_from=now + timedelta(days=1),
)
quote = quote_recharge_points(Decimal("8.88"), currency="cny", at=now)
self.assertEqual(quote.exchange_rate, current)
self.assertEqual(quote.currency, "CNY")
self.assertEqual(quote.points_per_unit, Decimal("12.5000"))
self.assertEqual(quote.points_granted, 111)
self.assertEqual(calculate_points_granted("8.88", currency="CNY", at=now), 111)
class BillingServiceTests(TestCase):
def setUp(self):
suffix = uuid.uuid4().hex[:8]
self.user = get_user_model().objects.create_user(
username=f"charged-client-{suffix}",
email=f"charged-client-{suffix}@example.com",
password="password",
)
self.wallet = UserWallet.objects.create(user=self.user, points_balance=100)
self.api_key, _raw_key = ApiKey.create_for_user(self.user, name="server")
def create_recharge_order(
self,
*,
amount="20.00",
points_granted=200,
pay_method=RechargeOrder.PayMethod.WEIXIN,
) -> RechargeOrder:
return RechargeOrder.objects.create(
user=self.user,
order_no=f"R{uuid.uuid4().hex[:12]}",
amount_money=Decimal(amount),
pay_method=pay_method,
exchange_rate=Decimal("10.0000"),
points_granted=points_granted,
)
def test_precharge_call_debits_wallet_and_writes_pending_call_and_consume_ledger(self):
charge = precharge_call(
user=self.user,
api_key=self.api_key,
operation_type=CallRecord.OperationType.TITLE,
alias="title-standard",
model_used="gpt-5.5",
resolution="1k",
prompt="Generate titles",
points_cost=12,
)
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 88)
self.assertEqual(charge.balance_after, 88)
self.assertEqual(charge.points_cost, 12)
call = charge.call_record
self.assertEqual(call.status, CallRecord.Status.PENDING)
self.assertEqual(call.user, self.user)
self.assertEqual(call.api_key, self.api_key)
self.assertEqual(call.alias, "title-standard")
self.assertEqual(call.resolution, "1K")
self.assertEqual(call.points_cost, 12)
ledger = charge.ledger_entry
self.assertEqual(ledger.change_type, PointsLedger.ChangeType.CONSUME)
self.assertEqual(ledger.points_delta, -12)
self.assertEqual(ledger.balance_after, 88)
self.assertEqual(ledger.ref_call, call)
def test_precharge_call_rejects_insufficient_points_without_call_or_ledger(self):
self.wallet.points_balance = 5
self.wallet.save(update_fields=("points_balance", "updated_at"))
with self.assertRaises(InsufficientPointsError) as context:
precharge_call(
user=self.user,
operation_type=CallRecord.OperationType.IMAGE,
alias="image-hd",
model_used="gpt-image-2",
points_cost=10,
)
self.assertEqual(context.exception.code, "insufficient_points")
self.assertEqual(context.exception.required_points, 10)
self.assertEqual(context.exception.available_points, 5)
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 5)
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
self.assertFalse(PointsLedger.objects.filter(user=self.user).exists())
def test_refund_call_points_restores_balance_marks_failed_and_is_idempotent(self):
charge = precharge_call(
user=self.user,
operation_type=CallRecord.OperationType.IMAGE,
alias="image-hd",
model_used="gpt-image-2",
points_cost=30,
)
refund = refund_call_points(
charge.call_record,
error_message="upstream timeout",
reason="provider failed",
)
self.wallet.refresh_from_db()
refund.call_record.refresh_from_db()
self.assertTrue(refund.refunded)
self.assertEqual(refund.points_refunded, 30)
self.assertEqual(refund.balance_after, 100)
self.assertEqual(self.wallet.points_balance, 100)
self.assertEqual(refund.call_record.status, CallRecord.Status.FAILED)
self.assertEqual(refund.call_record.error_message, "upstream timeout")
self.assertEqual(refund.ledger_entry.change_type, PointsLedger.ChangeType.REFUND)
self.assertEqual(refund.ledger_entry.points_delta, 30)
self.assertEqual(refund.ledger_entry.balance_after, 100)
second_refund = refund_call_points(charge.call_record, error_message="retry timeout")
self.wallet.refresh_from_db()
self.assertFalse(second_refund.refunded)
self.assertEqual(second_refund.points_refunded, 0)
self.assertEqual(self.wallet.points_balance, 100)
self.assertEqual(
PointsLedger.objects.filter(
ref_call=charge.call_record,
change_type=PointsLedger.ChangeType.REFUND,
).count(),
1,
)
def test_mark_call_success_does_not_change_balance_and_blocks_failure_refund(self):
charge = precharge_call(
user=self.user,
operation_type=CallRecord.OperationType.TITLE,
alias="title-standard",
model_used="gpt-5.5",
points_cost=8,
)
call = mark_call_success(
charge.call_record,
result_ref="https://cdn.example.test/result.txt",
result_summary="3 titles",
upstream_latency_ms=42,
)
self.wallet.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 92)
self.assertEqual(call.status, CallRecord.Status.SUCCESS)
self.assertEqual(call.result_ref, "https://cdn.example.test/result.txt")
self.assertEqual(call.result_summary, "3 titles")
self.assertEqual(call.upstream_latency_ms, 42)
with self.assertRaises(InvalidCallStateError):
refund_call_points(call, error_message="late failure")
def test_apply_recharge_payment_credits_wallet_writes_ledger_and_marks_paid(self):
order = self.create_recharge_order(amount="20.00", points_granted=200)
result = apply_recharge_payment(
RechargePayment(
order_no=order.order_no,
pay_method=RechargeOrder.PayMethod.WEIXIN,
amount=Decimal("20.00"),
transaction_id="wx-txn-001",
paid_at=timezone.now(),
)
)
self.wallet.refresh_from_db()
order.refresh_from_db()
self.assertTrue(result.applied)
self.assertEqual(result.points_granted, 200)
self.assertEqual(result.balance_after, 300)
self.assertEqual(self.wallet.points_balance, 300)
self.assertEqual(order.status, RechargeOrder.Status.PAID)
self.assertEqual(order.payment_txn_no, "wx-txn-001")
self.assertIsNotNone(order.paid_at)
ledger = PointsLedger.objects.get(ref_order_id=order.id)
self.assertEqual(ledger.change_type, PointsLedger.ChangeType.RECHARGE)
self.assertEqual(ledger.points_delta, 200)
self.assertEqual(ledger.balance_after, 300)
self.assertEqual(ledger.user, self.user)
def test_apply_recharge_payment_is_idempotent_for_duplicate_callback(self):
order = self.create_recharge_order(amount="20.00", points_granted=200)
payment = RechargePayment(
order_no=order.order_no,
pay_method=RechargeOrder.PayMethod.WEIXIN,
amount=Decimal("20.00"),
transaction_id="wx-txn-duplicate",
paid_at=timezone.now(),
)
first = apply_recharge_payment(payment)
second = apply_recharge_payment(payment)
self.wallet.refresh_from_db()
self.assertTrue(first.applied)
self.assertFalse(second.applied)
self.assertEqual(self.wallet.points_balance, 300)
self.assertEqual(
PointsLedger.objects.filter(
ref_order_id=order.id,
change_type=PointsLedger.ChangeType.RECHARGE,
).count(),
1,
)
def test_apply_recharge_payment_rejects_amount_mismatch_without_crediting(self):
order = self.create_recharge_order(amount="20.00", points_granted=200)
with self.assertRaises(RechargeAmountMismatchError):
apply_recharge_payment(
RechargePayment(
order_no=order.order_no,
pay_method=RechargeOrder.PayMethod.WEIXIN,
amount=Decimal("19.99"),
transaction_id="wx-txn-bad-amount",
paid_at=timezone.now(),
)
)
self.wallet.refresh_from_db()
order.refresh_from_db()
self.assertEqual(self.wallet.points_balance, 100)
self.assertEqual(order.status, RechargeOrder.Status.PENDING)
self.assertFalse(PointsLedger.objects.filter(ref_order_id=order.id).exists())
def test_query_and_apply_recharge_payment_uses_same_idempotent_path(self):
order = self.create_recharge_order(amount="30.00", points_granted=300)
def fake_query(queried_order):
return RechargePayment(
order_no=queried_order.order_no,
pay_method=queried_order.pay_method,
amount=queried_order.amount_money,
transaction_id="queried-txn-001",
paid_at=timezone.now(),
)
result = query_and_apply_recharge_payment(order.order_no, fake_query)
self.wallet.refresh_from_db()
order.refresh_from_db()
self.assertTrue(result.applied)
self.assertEqual(self.wallet.points_balance, 400)
self.assertEqual(order.status, RechargeOrder.Status.PAID)
self.assertEqual(order.payment_txn_no, "queried-txn-001")
class ConcurrentDebitTests(TransactionTestCase):
def setUp(self):
self.addCleanup(connections.close_all)
suffix = uuid.uuid4().hex[:8]
self.user = get_user_model().objects.create_user(
username=f"concurrent-client-{suffix}",
email=f"concurrent-client-{suffix}@example.com",
password="password",
)
UserWallet.objects.create(user=self.user, points_balance=30)
def test_concurrent_precharge_does_not_overspend_or_make_balance_negative(self):
worker_count = 2
barrier = threading.Barrier(worker_count)
results = []
errors = []
results_lock = threading.Lock()
def worker(index):
connections.close_all()
try:
for attempt in range(3):
try:
connections["default"].ensure_connection()
break
except OperationalError as exc:
connections.close_all()
if attempt == 2:
with results_lock:
errors.append(repr(exc))
barrier.abort()
return
time.sleep(0.5)
barrier.wait(timeout=15)
charge = precharge_call(
user=self.user,
operation_type=CallRecord.OperationType.TITLE,
alias="title-standard",
model_used="gpt-5.5",
points_cost=30,
prompt=f"worker {index}",
)
with results_lock:
results.append(("success", charge.balance_after))
except InsufficientPointsError as exc:
with results_lock:
results.append((exc.code, exc.available_points))
except Exception as exc: # pragma: no cover - surfaced through assertion below.
with results_lock:
errors.append(repr(exc))
finally:
connections.close_all()
threads = [threading.Thread(target=worker, args=(index,)) for index in range(worker_count)]
for thread in threads:
thread.start()
for thread in threads:
thread.join(timeout=30)
self.assertFalse(any(thread.is_alive() for thread in threads), "worker thread timed out")
self.assertEqual(errors, [])
success_count = sum(1 for status, _value in results if status == "success")
insufficient_count = sum(1 for status, _value in results if status == "insufficient_points")
self.assertEqual(success_count, 1)
self.assertEqual(insufficient_count, 1)
wallet = UserWallet.objects.get(user=self.user)
self.assertEqual(wallet.points_balance, 0)
self.assertEqual(
PointsLedger.objects.filter(
user=self.user,
change_type=PointsLedger.ChangeType.CONSUME,
).count(),
1,
)
self.assertEqual(CallRecord.objects.filter(user=self.user).count(), 1)
self.assertFalse(
PointsLedger.objects.filter(
user=self.user,
balance_after__lt=0,
).exists()
)