From be97a35360c104193033dca0dd9ca0bbd6d7f595 Mon Sep 17 00:00:00 2001 From: QiuSW <105186638@qq.com> Date: Thu, 2 Jul 2026 16:34:45 +0800 Subject: [PATCH] feat: add billing debit refund services --- apps/billing/services.py | 201 ++++++++++++++++++++++++++++++ apps/billing/tests.py | 263 ++++++++++++++++++++++++++++++++++++--- 2 files changed, 447 insertions(+), 17 deletions(-) create mode 100644 apps/billing/services.py diff --git a/apps/billing/services.py b/apps/billing/services.py new file mode 100644 index 0000000..b5c9591 --- /dev/null +++ b/apps/billing/services.py @@ -0,0 +1,201 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from django.db import transaction + +from apps.users.models import UserWallet + +from .models import CallRecord, PointsLedger, normalize_resolution + + +class BillingOperationError(RuntimeError): + code = "billing_operation_error" + + +class InsufficientPointsError(BillingOperationError): + code = "insufficient_points" + + def __init__(self, *, required_points: int, available_points: int): + self.required_points = required_points + self.available_points = available_points + super().__init__("Insufficient points. Please recharge before calling this API.") + + +class InvalidCallStateError(BillingOperationError): + code = "invalid_call_state" + + def __init__(self, message: str): + super().__init__(message) + + +@dataclass(frozen=True) +class CallCharge: + call_record: CallRecord + ledger_entry: PointsLedger + points_cost: int + balance_after: int + + +@dataclass(frozen=True) +class RefundResult: + call_record: CallRecord + ledger_entry: PointsLedger | None + points_refunded: int + balance_after: int + refunded: bool + + +def _validate_positive_points(points: int) -> int: + if isinstance(points, bool) or not isinstance(points, int) or points <= 0: + raise ValueError("points must be a positive integer") + return points + + +def _locked_wallet_for_user(user) -> UserWallet: + wallet, _created = UserWallet.objects.select_for_update().get_or_create(user=user) + return wallet + + +def precharge_call( + *, + user, + points_cost: int, + operation_type: str, + alias: str, + model_used: str = "", + resolution: str | None = None, + api_key=None, + prompt: str = "", +) -> CallCharge: + points_cost = _validate_positive_points(points_cost) + normalized_alias = str(alias or "").strip() + normalized_resolution = normalize_resolution(resolution) + + with transaction.atomic(): + wallet = _locked_wallet_for_user(user) + if wallet.points_balance < points_cost: + raise InsufficientPointsError( + required_points=points_cost, + available_points=wallet.points_balance, + ) + + wallet.points_balance -= points_cost + wallet.save(update_fields=("points_balance", "updated_at")) + + call_record = CallRecord.objects.create( + user=user, + api_key=api_key, + operation_type=operation_type, + alias=normalized_alias, + model_used=str(model_used or "").strip(), + resolution=normalized_resolution, + prompt=str(prompt or ""), + points_cost=points_cost, + status=CallRecord.Status.PENDING, + ) + ledger_entry = PointsLedger.objects.create( + user=user, + change_type=PointsLedger.ChangeType.CONSUME, + points_delta=-points_cost, + balance_after=wallet.points_balance, + ref_call=call_record, + ) + + return CallCharge( + call_record=call_record, + ledger_entry=ledger_entry, + points_cost=points_cost, + balance_after=ledger_entry.balance_after, + ) + + +def mark_call_success( + call_record: CallRecord, + *, + result_ref: str = "", + result_summary: str = "", + upstream_latency_ms: int | None = None, +) -> CallRecord: + with transaction.atomic(): + locked_call = CallRecord.objects.select_for_update().get(pk=call_record.pk) + if locked_call.status == CallRecord.Status.FAILED: + raise InvalidCallStateError("Cannot mark a failed call as successful.") + + locked_call.status = CallRecord.Status.SUCCESS + locked_call.result_ref = str(result_ref or "") + locked_call.result_summary = str(result_summary or "") + locked_call.upstream_latency_ms = upstream_latency_ms + locked_call.error_message = "" + locked_call.save( + update_fields=( + "status", + "result_ref", + "result_summary", + "upstream_latency_ms", + "error_message", + "updated_at", + ) + ) + + return locked_call + + +def refund_call_points( + call_record: CallRecord, + *, + error_message: str = "", + reason: str = "", +) -> RefundResult: + with transaction.atomic(): + locked_call = CallRecord.objects.select_for_update().get(pk=call_record.pk) + if locked_call.status == CallRecord.Status.SUCCESS: + raise InvalidCallStateError("Cannot refund a successful call via failure refund.") + + existing_refund = ( + PointsLedger.objects.filter( + ref_call=locked_call, + change_type=PointsLedger.ChangeType.REFUND, + ) + .order_by("id") + .first() + ) + if existing_refund is not None: + if locked_call.status != CallRecord.Status.FAILED: + locked_call.status = CallRecord.Status.FAILED + locked_call.error_message = str(error_message or locked_call.error_message or "") + locked_call.save(update_fields=("status", "error_message", "updated_at")) + return RefundResult( + call_record=locked_call, + ledger_entry=existing_refund, + points_refunded=0, + balance_after=existing_refund.balance_after, + refunded=False, + ) + + points_cost = _validate_positive_points(locked_call.points_cost) + wallet = _locked_wallet_for_user(locked_call.user) + wallet.points_balance += points_cost + wallet.save(update_fields=("points_balance", "updated_at")) + + locked_call.status = CallRecord.Status.FAILED + locked_call.error_message = str(error_message or "") + locked_call.save(update_fields=("status", "error_message", "updated_at")) + + ledger_reason = str(reason or error_message or "Call failed; refund precharged points.").strip() + ledger_entry = PointsLedger.objects.create( + user=locked_call.user, + change_type=PointsLedger.ChangeType.REFUND, + points_delta=points_cost, + balance_after=wallet.points_balance, + ref_call=locked_call, + reason=ledger_reason, + ) + + return RefundResult( + call_record=locked_call, + ledger_entry=ledger_entry, + points_refunded=points_cost, + balance_after=ledger_entry.balance_after, + refunded=True, + ) diff --git a/apps/billing/tests.py b/apps/billing/tests.py index 5c7e9cd..cd543d4 100644 --- a/apps/billing/tests.py +++ b/apps/billing/tests.py @@ -1,3 +1,6 @@ +import threading +import time +import uuid from datetime import timedelta from decimal import Decimal @@ -5,8 +8,8 @@ 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, transaction -from django.test import TestCase, override_settings +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 @@ -18,6 +21,13 @@ from apps.billing.pricing import ( get_pricing_rule, quote_recharge_points, ) +from apps.billing.services import ( + InsufficientPointsError, + InvalidCallStateError, + mark_call_success, + precharge_call, + refund_call_points, +) from apps.users.models import ApiKey, UserWallet TEST_ENCRYPTION_KEY = Fernet.generate_key().decode("ascii") @@ -25,9 +35,10 @@ 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="client", - email="client@example.com", + username=f"client-{suffix}", + email=f"client-{suffix}@example.com", password="password", ) @@ -37,8 +48,8 @@ class BillingCoreModelTests(TestCase): self.assertEqual(wallet.points_balance, 0) other_user = get_user_model().objects.create_user( - username="negative", - email="negative@example.com", + username=f"negative-{uuid.uuid4().hex[:8]}", + email=f"negative-{uuid.uuid4().hex[:8]}@example.com", password="password", ) with self.assertRaises(IntegrityError): @@ -142,19 +153,22 @@ class BillingCoreModelTests(TestCase): @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="gpt-5.5-a", + model=f"gpt-5.5-a-{suffix}", capabilities=["text"], ) self.replacement_text_model = self.create_ai_model( name="Text model B", - model="gpt-5.5-b", + model=f"gpt-5.5-b-{suffix}", capabilities=["text"], ) self.alias = ModelAlias.objects.create( operation_type=ModelAlias.OperationType.TITLE, - alias="title-standard", + alias=self.title_alias, ai_model=self.text_model, ) @@ -173,13 +187,13 @@ class PricingCalculationTests(TestCase): def test_exact_resolution_rule_overrides_alias_default(self): PricingRule.objects.create( operation_type=CallRecord.OperationType.IMAGE, - alias="image-hd", + alias=self.image_alias, resolution="", points_cost=10, ) PricingRule.objects.create( operation_type=CallRecord.OperationType.IMAGE, - alias="image-hd", + alias=self.image_alias, resolution="1k", points_cost=15, ) @@ -187,7 +201,7 @@ class PricingCalculationTests(TestCase): self.assertEqual( calculate_points_cost( CallRecord.OperationType.IMAGE, - "image-hd", + self.image_alias, "1K", ), 15, @@ -195,7 +209,7 @@ class PricingCalculationTests(TestCase): self.assertEqual( calculate_points_cost( CallRecord.OperationType.IMAGE, - "image-hd", + self.image_alias, "2K", ), 10, @@ -203,7 +217,7 @@ class PricingCalculationTests(TestCase): exact_rule = get_pricing_rule( CallRecord.OperationType.IMAGE, - "image-hd", + self.image_alias, "1k", ) self.assertEqual(exact_rule.resolution, "1K") @@ -211,14 +225,14 @@ class PricingCalculationTests(TestCase): def test_pricing_is_bound_to_alias_not_underlying_model(self): PricingRule.objects.create( operation_type=CallRecord.OperationType.TITLE, - alias="title-standard", + alias=self.title_alias, points_cost=2, ) self.assertEqual( calculate_points_cost( CallRecord.OperationType.TITLE, - "title-standard", + self.title_alias, "1K", ), 2, @@ -230,7 +244,7 @@ class PricingCalculationTests(TestCase): self.assertEqual( calculate_points_cost( CallRecord.OperationType.TITLE, - "title-standard", + self.title_alias, "1K", ), 2, @@ -277,3 +291,218 @@ class PricingCalculationTests(TestCase): 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 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") + + +class ConcurrentDebitTests(TransactionTestCase): + def setUp(self): + 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() + )