feat: add billing debit refund services
This commit is contained in:
@@ -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,
|
||||||
|
)
|
||||||
+246
-17
@@ -1,3 +1,6 @@
|
|||||||
|
import threading
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
from datetime import timedelta
|
from datetime import timedelta
|
||||||
from decimal import Decimal
|
from decimal import Decimal
|
||||||
|
|
||||||
@@ -5,8 +8,8 @@ from cryptography.fernet import Fernet
|
|||||||
from django.contrib import admin
|
from django.contrib import admin
|
||||||
from django.contrib.auth import get_user_model
|
from django.contrib.auth import get_user_model
|
||||||
from django.core.exceptions import ValidationError
|
from django.core.exceptions import ValidationError
|
||||||
from django.db import IntegrityError, transaction
|
from django.db import IntegrityError, OperationalError, connections, transaction
|
||||||
from django.test import TestCase, override_settings
|
from django.test import TestCase, TransactionTestCase, override_settings
|
||||||
from django.utils import timezone
|
from django.utils import timezone
|
||||||
|
|
||||||
from apps.ai.models import AiModel, ModelAlias
|
from apps.ai.models import AiModel, ModelAlias
|
||||||
@@ -18,6 +21,13 @@ from apps.billing.pricing import (
|
|||||||
get_pricing_rule,
|
get_pricing_rule,
|
||||||
quote_recharge_points,
|
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
|
from apps.users.models import ApiKey, UserWallet
|
||||||
|
|
||||||
TEST_ENCRYPTION_KEY = Fernet.generate_key().decode("ascii")
|
TEST_ENCRYPTION_KEY = Fernet.generate_key().decode("ascii")
|
||||||
@@ -25,9 +35,10 @@ TEST_ENCRYPTION_KEY = Fernet.generate_key().decode("ascii")
|
|||||||
|
|
||||||
class BillingCoreModelTests(TestCase):
|
class BillingCoreModelTests(TestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
|
suffix = uuid.uuid4().hex[:8]
|
||||||
self.user = get_user_model().objects.create_user(
|
self.user = get_user_model().objects.create_user(
|
||||||
username="client",
|
username=f"client-{suffix}",
|
||||||
email="client@example.com",
|
email=f"client-{suffix}@example.com",
|
||||||
password="password",
|
password="password",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -37,8 +48,8 @@ class BillingCoreModelTests(TestCase):
|
|||||||
self.assertEqual(wallet.points_balance, 0)
|
self.assertEqual(wallet.points_balance, 0)
|
||||||
|
|
||||||
other_user = get_user_model().objects.create_user(
|
other_user = get_user_model().objects.create_user(
|
||||||
username="negative",
|
username=f"negative-{uuid.uuid4().hex[:8]}",
|
||||||
email="negative@example.com",
|
email=f"negative-{uuid.uuid4().hex[:8]}@example.com",
|
||||||
password="password",
|
password="password",
|
||||||
)
|
)
|
||||||
with self.assertRaises(IntegrityError):
|
with self.assertRaises(IntegrityError):
|
||||||
@@ -142,19 +153,22 @@ class BillingCoreModelTests(TestCase):
|
|||||||
@override_settings(AI_KEY_ENCRYPTION_KEY=TEST_ENCRYPTION_KEY)
|
@override_settings(AI_KEY_ENCRYPTION_KEY=TEST_ENCRYPTION_KEY)
|
||||||
class PricingCalculationTests(TestCase):
|
class PricingCalculationTests(TestCase):
|
||||||
def setUp(self):
|
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(
|
self.text_model = self.create_ai_model(
|
||||||
name="Text model A",
|
name="Text model A",
|
||||||
model="gpt-5.5-a",
|
model=f"gpt-5.5-a-{suffix}",
|
||||||
capabilities=["text"],
|
capabilities=["text"],
|
||||||
)
|
)
|
||||||
self.replacement_text_model = self.create_ai_model(
|
self.replacement_text_model = self.create_ai_model(
|
||||||
name="Text model B",
|
name="Text model B",
|
||||||
model="gpt-5.5-b",
|
model=f"gpt-5.5-b-{suffix}",
|
||||||
capabilities=["text"],
|
capabilities=["text"],
|
||||||
)
|
)
|
||||||
self.alias = ModelAlias.objects.create(
|
self.alias = ModelAlias.objects.create(
|
||||||
operation_type=ModelAlias.OperationType.TITLE,
|
operation_type=ModelAlias.OperationType.TITLE,
|
||||||
alias="title-standard",
|
alias=self.title_alias,
|
||||||
ai_model=self.text_model,
|
ai_model=self.text_model,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -173,13 +187,13 @@ class PricingCalculationTests(TestCase):
|
|||||||
def test_exact_resolution_rule_overrides_alias_default(self):
|
def test_exact_resolution_rule_overrides_alias_default(self):
|
||||||
PricingRule.objects.create(
|
PricingRule.objects.create(
|
||||||
operation_type=CallRecord.OperationType.IMAGE,
|
operation_type=CallRecord.OperationType.IMAGE,
|
||||||
alias="image-hd",
|
alias=self.image_alias,
|
||||||
resolution="",
|
resolution="",
|
||||||
points_cost=10,
|
points_cost=10,
|
||||||
)
|
)
|
||||||
PricingRule.objects.create(
|
PricingRule.objects.create(
|
||||||
operation_type=CallRecord.OperationType.IMAGE,
|
operation_type=CallRecord.OperationType.IMAGE,
|
||||||
alias="image-hd",
|
alias=self.image_alias,
|
||||||
resolution="1k",
|
resolution="1k",
|
||||||
points_cost=15,
|
points_cost=15,
|
||||||
)
|
)
|
||||||
@@ -187,7 +201,7 @@ class PricingCalculationTests(TestCase):
|
|||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
calculate_points_cost(
|
calculate_points_cost(
|
||||||
CallRecord.OperationType.IMAGE,
|
CallRecord.OperationType.IMAGE,
|
||||||
"image-hd",
|
self.image_alias,
|
||||||
"1K",
|
"1K",
|
||||||
),
|
),
|
||||||
15,
|
15,
|
||||||
@@ -195,7 +209,7 @@ class PricingCalculationTests(TestCase):
|
|||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
calculate_points_cost(
|
calculate_points_cost(
|
||||||
CallRecord.OperationType.IMAGE,
|
CallRecord.OperationType.IMAGE,
|
||||||
"image-hd",
|
self.image_alias,
|
||||||
"2K",
|
"2K",
|
||||||
),
|
),
|
||||||
10,
|
10,
|
||||||
@@ -203,7 +217,7 @@ class PricingCalculationTests(TestCase):
|
|||||||
|
|
||||||
exact_rule = get_pricing_rule(
|
exact_rule = get_pricing_rule(
|
||||||
CallRecord.OperationType.IMAGE,
|
CallRecord.OperationType.IMAGE,
|
||||||
"image-hd",
|
self.image_alias,
|
||||||
"1k",
|
"1k",
|
||||||
)
|
)
|
||||||
self.assertEqual(exact_rule.resolution, "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):
|
def test_pricing_is_bound_to_alias_not_underlying_model(self):
|
||||||
PricingRule.objects.create(
|
PricingRule.objects.create(
|
||||||
operation_type=CallRecord.OperationType.TITLE,
|
operation_type=CallRecord.OperationType.TITLE,
|
||||||
alias="title-standard",
|
alias=self.title_alias,
|
||||||
points_cost=2,
|
points_cost=2,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
calculate_points_cost(
|
calculate_points_cost(
|
||||||
CallRecord.OperationType.TITLE,
|
CallRecord.OperationType.TITLE,
|
||||||
"title-standard",
|
self.title_alias,
|
||||||
"1K",
|
"1K",
|
||||||
),
|
),
|
||||||
2,
|
2,
|
||||||
@@ -230,7 +244,7 @@ class PricingCalculationTests(TestCase):
|
|||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
calculate_points_cost(
|
calculate_points_cost(
|
||||||
CallRecord.OperationType.TITLE,
|
CallRecord.OperationType.TITLE,
|
||||||
"title-standard",
|
self.title_alias,
|
||||||
"1K",
|
"1K",
|
||||||
),
|
),
|
||||||
2,
|
2,
|
||||||
@@ -277,3 +291,218 @@ class PricingCalculationTests(TestCase):
|
|||||||
self.assertEqual(quote.points_per_unit, Decimal("12.5000"))
|
self.assertEqual(quote.points_per_unit, Decimal("12.5000"))
|
||||||
self.assertEqual(quote.points_granted, 111)
|
self.assertEqual(quote.points_granted, 111)
|
||||||
self.assertEqual(calculate_points_granted("8.88", currency="CNY", at=now), 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()
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user