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, )