@@ -194,4 +194,76 @@ PATH_PREFETCH_MAP = { *_gen_only('payment_plan__plan', 'uid', 'price'), ), }, + '/api/v1/payments/restore-subscription': { + 'select': ( + 'host_account', + 'host_account__company_companyipwhitelist', + 'business_account', + 'business_account__parent_company', + 'business_account__parent_company__company_companyipwhitelist', + 'payment_plan', + 'payment_plan__plan', + ), + 'prefetch': ( + Prefetch( + 'payment_plan__methods', + queryset=PaymentMethod.objects.filter(active=True) + .order_by('-primary', '-created_at') + .only('uid', 'created_at', 'payment_method_id', 'primary', 'active', 'user_plan_info_id'), + to_attr='active_methods', + ), + ), + 'only': ( + 'uid', + 'email', + 'is_staff', + 'is_superuser', + *_gen_only('host_account', 'uid'), + *_gen_only('host_account__company_companyipwhitelist', 'uid', 'is_enabled'), + *_gen_only('business_account', 'account_privileges'), + *_gen_only('business_account__parent_company', 'uid'), + *_gen_only( + 'business_account__parent_company__company_companyipwhitelist', + 'uid', + 'is_enabled', + ), + *_gen_only('payment_plan', 'uid'), + *_gen_only( + 'payment_plan__plan', + 'uid', + 'price', + 'tokens_per_plan', + ), + ), + }, + '/api/v1/payments/restore-subscription/blocked': { + 'select': ( + 'host_account', + 'host_account__company_companyipwhitelist', + 'business_account', + 'business_account__parent_company', + 'business_account__parent_company__company_companyipwhitelist', + 'payment_plan', + ), + 'only': ( + 'uid', + 'email', + 'is_staff', + 'is_superuser', + *_gen_only('host_account', 'uid'), + *_gen_only('host_account__company_companyipwhitelist', 'uid', 'is_enabled'), + *_gen_only('business_account', 'account_privileges'), + *_gen_only('business_account__parent_company', 'uid'), + *_gen_only( + 'business_account__parent_company__company_companyipwhitelist', + 'uid', + 'is_enabled', + ), + *_gen_only( + 'payment_plan', + 'uid', + 'recovery_locked_at', + ), + ), + }, } @@ -97,24 +97,7 @@ def invalid_password_error_handler(request, exc: InvalidPassword): return api.create_response(request, {'message': _('Wrong password')}, status=401) -urlpatterns = ( - [ - path('admin/', admin.site.urls), - path('api/v1/ml_models/', include('ml_model.urls', namespace='ml_model')), - path('api/v1/auth/', include('authentication.urls')), - path('api/v1/payments/', include('payments.urls')), - path('api/v1/reports/', include('reports.urls')), - path('api/v1/chats/', include('tools.chats.urls')), - path('api/v1/media/', include('tools.media.urls')), - path('api/v1/public/', include('tools.public_api.urls')), - path('api/v1/api/', api.urls), - path('api/v1/', compatibility_api.urls), - path('api/v1/v2/', compatibility_api_v2.urls), - path('public/', public_api.urls), - ] - + static(settings.STATIC_URL, document_root=settings.STATIC_ROOT) - + public_urlpatterns -) +urlpatterns = [] if settings.DEBUG: urlpatterns += [ @@ -137,3 +120,22 @@ if settings.DEBUG: compatibility_api.docs_url = '/docs' compatibility_api_v2.docs_url = '/docs' public_api.docs_url = '/docs' + +urlpatterns += ( + [ + path('admin/', admin.site.urls), + path('api/v1/ml_models/', include('ml_model.urls', namespace='ml_model')), + path('api/v1/auth/', include('authentication.urls')), + path('api/v1/payments/', include('payments.urls')), + path('api/v1/reports/', include('reports.urls')), + path('api/v1/chats/', include('tools.chats.urls')), + path('api/v1/media/', include('tools.media.urls')), + path('api/v1/public/', include('tools.public_api.urls')), + path('api/v1/api/', api.urls), + path('api/v1/', compatibility_api.urls), + path('api/v1/v2/', compatibility_api_v2.urls), + path('public/', public_api.urls), + ] + + static(settings.STATIC_URL, document_root=settings.STATIC_ROOT) + + public_urlpatterns +) \ No newline at end of file @@ -760,11 +760,8 @@ msgid "Prompt is too long. Maximum length is %(max_length)s characters." msgstr "Промпт слишком длинный. Максимальная длина — %(max_length)s символов." #: ml_model/exceptions.py:161 -msgid "" -"Service is currently unavailable due to high demand. Please try again later" -msgstr "" -"Сервис временно недоступен из-за высокой нагрузки. Пожалуйста, попробуйте " -"позже" +msgid "Service is temporarily unavailable. Please try again later" +msgstr "Сервис временно недоступен. Пожалуйста, попробуйте позже" #: ml_model/exceptions.py:169 #, python-format @@ -1368,6 +1365,18 @@ msgstr "У вас нет активной подписки для отмены" msgid "The recurring payment is successfully cancelled" msgstr "Автоплатежи успешно отключены" +#: payments/exceptions/subscription_recovery.py +msgid "Active payment method not found" +msgstr "Активный способ оплаты не найден" + +#: payments/exceptions/subscription_recovery.py +msgid "Subscription recovery is already in progress" +msgstr "Восстановление подписки уже выполняется" + +#: payments/routes/v1.py +msgid "Payment could not be completed, please try again later" +msgstr "Не удалось провести платёж, пожалуйста, попробуйте позже" + #: payments/routes/v1.py:144 msgid "Expenses" msgstr "Затраты" @@ -1,19 +1,13 @@ import time -import requests - from datetime import timedelta from decimal import Decimal from io import BytesIO from typing import Any +import requests from django.core.files import File from messages.models import Message -from ml_model.exceptions import ( - ImageContentNotFound, - GenerationException, - RequestBlocked, -) from ml_model.services.base import SimpleService from ml_model.tasks import replicate_run from payments.exceptions.insufficient_balance import InsufficientBalance @@ -51,12 +45,12 @@ class Grok_Image(SimpleService): raise InsufficientBalance(balance, self.TOKENS_COST) callback_data = dict( { - 'prompt': f"{self.translate_prompt(input_message.content)}\n{self.OPTIMIZATION_PROMPT}", + 'prompt': f'{self.translate_prompt(input_message.content)}\n{self.OPTIMIZATION_PROMPT}', **input_message.info, } ) if image := input_message.file: - callback_data.update({'image': image.url}) + callback_data.update({'image': image}) start_time = time.time() images = replicate_run('xai/grok-imagine-image', callback_data) process_time = timedelta(seconds=(time.time() - start_time)) @@ -68,14 +68,14 @@ class Grok_Image_Ultra(SimpleService): file_bytes = input_message.file.read() kind = filetype.guess(file_bytes[:20]) extension = kind.extension - if extension.upper() not in (extensions := ['JPG', 'JPEG', 'PNG', 'WEBP']): + if extension.upper() not in (extensions := ['JPG', 'JPEG', 'JFIF', 'PNG', 'WEBP']): raise FileExtensionNotSupported(extensions) file_width, file_height = get_image_dimensions(BytesIO(file_bytes)) if file_width and file_height: input_mp = math.ceil((file_width * file_height) / 1_000_000) else: input_mp = 1 - callback_data.update({'image': input_message.file.url}) + callback_data.update({'image': input_message.file}) input_message.file.close() if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < ( @@ -4,6 +4,7 @@ from typing import Any import httpx from django.conf import settings +from ml_model.exceptions import OpenAIResponseError, ServiceTemporaryUnavailableError from poller.models import Proxy type OpenAIEvent = dict[str, Any] @@ -26,6 +27,8 @@ class OpenAIStreamMixin: input_tokens, output_tokens = yield from self._stream_request( client, 'POST', json=payload, state=state ) + except ServiceTemporaryUnavailableError: + raise except Exception: if not state['response_id']: raise @@ -95,7 +98,9 @@ class OpenAIStreamMixin: return self._get_response_id(event) case 'response.output_text.delta': return self._get_delta(event) - case 'response.completed' | 'response.incomplete' | 'response.failed': + case 'response.failed': + raise ServiceTemporaryUnavailableError from self._get_response_error(event) + case 'response.completed' | 'response.incomplete': return self._get_usage(event) case _: return '' @@ -109,3 +114,11 @@ class OpenAIStreamMixin: def _get_response_id(self, event: OpenAIEvent) -> str: return event.get('response', {}).get('id', '') + + def _get_response_error(self, event: OpenAIEvent) -> OpenAIResponseError: + error = event.get('error') or (event.get('response') or {}).get('error') or {} + return OpenAIResponseError( + event_type=event.get('type', 'unknown'), + code=error.get('code', 'unknown'), + message=error.get('message', 'Unknown OpenAI response error'), + ) @@ -162,6 +162,23 @@ class ServiceHighDemandError(Exception): return _('Service is currently unavailable due to high demand. Please try again later') +class ServiceTemporaryUnavailableError(Exception): + def __str__(self) -> str: + return _('Service is temporarily unavailable. Please try again later') + + +class OpenAIResponseError(Exception): + def __init__(self, event_type: str, code: str, message: str) -> None: + self.event_type = event_type + self.code = code + self.message = message + + def __str__(self) -> str: + return ( + f'OpenAI streaming error: event_type={self.event_type}, code={self.code}, message={self.message}' + ) + + class PaidPlanRequiredError(Exception): def __init__(self, feature: str) -> None: self.feature = feature @@ -6,6 +6,7 @@ import time # import uuid from io import BytesIO +from pathlib import PurePosixPath from typing import IO, Any, Dict import deepl @@ -15,6 +16,8 @@ import replicate import requests from celery import shared_task from deepl.translator import TextResult +from django.core.files.storage import Storage +from django.db.models.fields.files import FieldFile from django.utils.translation import gettext as _ from replicate.exceptions import ModelError from requests import Response @@ -28,6 +31,7 @@ from ml_model.exceptions import ( DeploymentDisabled, ExceededContextLengthError, FaceNotFoundError, + FileExtensionNotSupported, GenerationException, ImageAnalysisError, ImageContentNotFound, @@ -36,6 +40,7 @@ from ml_model.exceptions import ( ModelTimeoutError, PredictionInterruptedError, RequestBlocked, + ServiceHighDemandError, ) from poller.models import Proxy @@ -115,10 +120,37 @@ def transcript_audio(payload: dict[str, Any]): ) +def _prepare_replicate_image(image: FieldFile) -> tuple[str, tuple[Storage, str] | None]: + if PurePosixPath(image.name).suffix.lower() != '.jfif': + return image.url, None + + # JFIF already contains JPEG data, so copy it without decoding or re-encoding. + image.open('rb') + image.seek(0) + jpeg_name = str(PurePosixPath(image.name).with_suffix('.jpg')) + try: + saved_name = image.storage.save(jpeg_name, image) + finally: + image.close() + + try: + image_url = image.storage.url(saved_name) + except Exception: + image.storage.delete(saved_name) + + raise + + return image_url, (image.storage, saved_name) + + @shared_task def replicate_run(callback_url: str, payload: dict[str, Any]): replicate_client = replicate.Client(settings.REPLICATE_API_KEY) + temporary_file = None try: + if isinstance(image := payload.get('image'), FieldFile): + payload['image'], temporary_file = _prepare_replicate_image(image) + return replicate_client.run( ref=callback_url, input=payload, @@ -126,6 +158,14 @@ def replicate_run(callback_url: str, payload: dict[str, Any]): except ModelError as exc: prediction_error = getattr(getattr(exc, 'prediction', None), 'error', '') or '' error_text = str(exc) + if 'ModelRateLimitError' in error_text or 'E003' in error_text: + raise ServiceHighDemandError from exc + if ( + 'Music upload failed' in error_text + and 'audio format' in error_text + and 'is not supported' in error_text + ): + raise FileExtensionNotSupported(('MP3', 'WAV')) from exc if any(error in error_text for error in ('E005', 'E006', 'sexual', 'NSFW')): raise RequestBlocked if 'PA' in error_text: @@ -147,6 +187,10 @@ def replicate_run(callback_url: str, payload: dict[str, Any]): if 'PROMPT_TOO_LONG' in error_text: raise ExceededContextLengthError raise GenerationException from exc + finally: + if temporary_file: + storage, file_name = temporary_file + storage.delete(file_name) @shared_task @@ -0,0 +1,11 @@ +from django.utils.translation import gettext as _ + + +class ActivePaymentMethodNotFound(Exception): + def __str__(self) -> str: + return _('Active payment method not found') + + +class DuplicateRecoveryAttempt(Exception): + def __str__(self) -> str: + return _('Subscription recovery is already in progress') @@ -0,0 +1,28 @@ +from django.db import migrations, models + + +class Migration(migrations.Migration): + dependencies = [ + ('payments', '0034_add_flux_payment_features'), + ] + + operations = [ + migrations.AddField( + model_name='paymentplanuserinfo', + name='last_recovery_payment_id', + field=models.UUIDField( + blank=True, + null=True, + verbose_name='Last recovery payment id', + ), + ), + migrations.AddField( + model_name='paymentplanuserinfo', + name='recovery_locked_at', + field=models.DateTimeField( + blank=True, + null=True, + verbose_name='Recovery locked at', + ), + ), + ] @@ -1,7 +1,6 @@ from datetime import datetime from django.contrib.auth import get_user_model -from django.contrib.postgres.fields import ArrayField from django.db import models from django.utils.translation import gettext_lazy as _ @@ -50,6 +49,18 @@ class PaymentPlanUserInfo(BaseModel): ) last_payment_at = models.DateField(verbose_name=_('Last payment at')) next_payment_at = models.DateTimeField(blank=True, null=True, verbose_name=_('Next payment at')) + # FIXME: вынести в абстракцию попыток + last_recovery_payment_id = models.UUIDField( + blank=True, + null=True, + verbose_name=_('Last recovery payment id'), + ) + # FIXME: вынести в абстракцию попыток + recovery_locked_at = models.DateTimeField( + blank=True, + null=True, + verbose_name=_('Recovery locked at'), + ) current_token_balance = models.DecimalField( max_digits=100, decimal_places=10, verbose_name=_('Current balance') ) @@ -1,31 +1,26 @@ -import orjson - import calendar import logging from collections import defaultdict from datetime import date, timedelta from decimal import Decimal -from django.db.models.aggregates import Count -from django.utils.translation import gettext_lazy as _ - +import orjson from dateutil.relativedelta import relativedelta from django.db.models import CharField, F, Func, Prefetch, Q, Sum, Value +from django.db.models.aggregates import Count from django.db.models.functions import Round, TruncDay, TruncMonth, TruncYear from django.utils.translation import gettext as _ from ninja import Query, Router from ninja.errors import HttpError +from authentication.exceptions.business_host_exceptions.access_denied import AccessDenied from authentication.models import CustomUserModel from authentication.security import SyncAuthBearer -from authentication.exceptions.business_host_exceptions.access_denied import AccessDenied -from payments.models import ( - Invoice, - Payment, - PaymentPlan, - PaymentPlanFeature, - PaymentMethod, +from payments.exceptions.subscription_recovery import ( + ActivePaymentMethodNotFound, + DuplicateRecoveryAttempt, ) +from payments.models import Invoice, Payment, PaymentMethod, PaymentPlan, PaymentPlanFeature from payments.schema import UserBalance from payments.schemas import ( ExpensesParamsSchema, @@ -33,11 +28,13 @@ from payments.schemas import ( NewSubscriptionSchema, PaymentLinkSchema, PaymentPlanSchema, + RestoreSubscriptionLockSchema, + RestoreSubscriptionSchema, ) from payments.selectors.payment_plan_selector import PaymentPlanSelector from payments.services.payment_method_service import PaymentMethodService -from payments.typing import IntervalStrategyEnum, SourceStrategyEnum from payments.services.payment_service import PaymentService +from payments.typing import IntervalStrategyEnum, SourceStrategyEnum router = Router(auth=SyncAuthBearer(), tags=['payments']) @@ -114,6 +111,44 @@ def revoke_recurring_payment(request): return 200, {'detail': _('The recurring payment is successfully cancelled')} +@router.post( + 'restore-subscription', + tags=['payments/restore-subscription'], + response=RestoreSubscriptionSchema, +) +def restore_subscription(request): + payment_service = PaymentService(request.auth) + try: + payment = payment_service.restore_subscription() + except ActivePaymentMethodNotFound as exc: + raise HttpError(400, str(exc)) from exc + except DuplicateRecoveryAttempt as exc: + raise HttpError(403, str(exc)) from exc + except Exception as exc: + logger.exception('Subscription recovery failed: email=%s', request.auth.email) + raise HttpError(500, _('Payment could not be completed, please try again later')) from exc + + try: + if result := payment_service.yookassa_payment_polling(payment): + return RestoreSubscriptionSchema(**result) + if result := payment_service.db_payment_polling(payment.id): + return RestoreSubscriptionSchema(**result) + return RestoreSubscriptionSchema(ok=True, external=False, internal=False) + finally: + payment_service.release_subscription_recovery_lock() + + +@router.get( + 'restore-subscription/blocked', + tags=['payments/restore-subscription'], + response=RestoreSubscriptionLockSchema, +) +def get_restore_subscription_blocked(request): + return RestoreSubscriptionLockSchema( + blocked=PaymentService(request.auth).is_subscription_recovery_blocked(), + ) + + @router.get('expenses', tags=['payments/expenses'], response=list[ExpensesSchema]) def list_expenses(request, data: ExpensesParamsSchema = Query(...)): try: @@ -1,6 +1,7 @@ +import hashlib import logging -from datetime import timedelta - +import time +from datetime import datetime, timedelta from decimal import Decimal from uuid import UUID, uuid4 @@ -13,8 +14,13 @@ from yookassa.domain.response import PaymentResponse as YookassaPaymentResponse from authentication.models import CustomUserModel from authentication.services.email_service import EmailService +from payments.exceptions.subscription_recovery import ( + ActivePaymentMethodNotFound, + DuplicateRecoveryAttempt, +) from payments.models.payment import Payment as PaymentModel from payments.models.payment_plan import PaymentPlan, PaymentPlanUserInfo +from payments.models.user_payment_method import PaymentMethod from payments.services.payment_method_service import PaymentMethodService from payments.services.referral_account import ReferralAccountService @@ -24,6 +30,8 @@ logger = logging.getLogger(__name__) class PaymentService: Configuration.account_id = settings.YOOKASSA_ACCOUNT_ID Configuration.secret_key = settings.YOOKASSA_SECRET_KEY + RECOVERY_LOCK_TTL = timedelta(minutes=2) + RECOVERY_BACKOFF_SECONDS = (0.25, 0.75, 1.0, 3.0) def __init__(self, user: CustomUserModel): self.user = user @@ -58,6 +66,108 @@ class PaymentService: ) return payment.confirmation.confirmation_url + def restore_subscription(self) -> YookassaPaymentResponse: + payment_method, generation = self._reserve_subscription_recovery() + plan = self.user.payment_plan.plan + try: + # FIXME: вынести сборку payload создания платежа в общий метод с create_payment_link + payment = YookassaPayment.create( + { + 'amount': {'value': f'{plan.price}', 'currency': 'RUB'}, + 'payment_method_id': payment_method.payment_method_id, + 'receipt': { + 'customer': {'email': self.user.email}, + 'items': [ + { + 'description': str(plan), + 'amount': {'value': f'{plan.price}', 'currency': 'RUB'}, + 'vat_code': 1, + 'quantity': '1', + } + ], + }, + 'description': str(self.user.uid), + 'capture': True, + 'metadata': { + 'plan_uid': str(plan.uid), + 'recovery': True, + }, + }, + idempotency_key=hashlib.sha256( + f'recovery:{self.user.uid}:{plan.uid}:{generation}'.encode() + ).hexdigest(), + ) + except Exception: + self.release_subscription_recovery_lock() + raise + logger.info( + 'Subscription recovery payment created: payment_id=%s email=%s method_uid=%s status=%s', + payment.id, + self.user.email, + payment_method.uid, + payment.status, + ) + return payment + + def yookassa_payment_polling(self, payment: YookassaPaymentResponse) -> dict[str, bool] | None: + current = payment + for delay in (0, *self.RECOVERY_BACKOFF_SECONDS): + if delay: + time.sleep(delay) + current = YookassaPayment.find_one(str(payment.id)) + if current.status == PaymentModel.SUCCEEDED: + return None + if current.status == PaymentModel.CANCELLED: + self._close_recovery_generation(current) + return {'ok': False, 'external': False, 'internal': False} + return {'ok': False, 'external': True, 'internal': False} + + def db_payment_polling(self, payment_id: UUID | str) -> dict[str, bool] | None: + for delay in (0, *self.RECOVERY_BACKOFF_SECONDS): + if delay: + time.sleep(delay) + if PaymentModel.objects.filter(uid=payment_id, status=PaymentModel.SUCCEEDED).exists(): + return None + return {'ok': False, 'external': False, 'internal': True} + + def is_subscription_recovery_blocked(self) -> bool: + locked_at = self.user.payment_plan.recovery_locked_at + return bool(locked_at) and not self._is_recovery_lock_stale(locked_at) + + def release_subscription_recovery_lock(self) -> None: + PaymentPlanUserInfo.objects.filter(pk=self.user.payment_plan.pk).update(recovery_locked_at=None) + + def _reserve_subscription_recovery(self) -> tuple[PaymentMethod, str]: + active_methods = self.user.payment_plan.active_methods + if not active_methods: + raise ActivePaymentMethodNotFound + + with transaction.atomic(): + info = PaymentPlanUserInfo.objects.select_for_update().get(pk=self.user.payment_plan.pk) + if info.recovery_locked_at and not self._is_recovery_lock_stale(info.recovery_locked_at): + logger.info('Duplicate subscription recovery skipped: email=%s', self.user.email) + raise DuplicateRecoveryAttempt + + info.recovery_locked_at = timezone.now() + info.save(update_fields=['recovery_locked_at']) + generation = str(info.last_recovery_payment_id or 'none') + + return active_methods[0], generation + + @classmethod + def _is_recovery_lock_stale(cls, locked_at: datetime) -> bool: + return locked_at <= timezone.now() - cls.RECOVERY_LOCK_TTL + + def _close_recovery_generation(self, payment: YookassaPaymentResponse) -> None: + metadata = payment.metadata or {} + if str(metadata.get('recovery', '')).lower() != 'true': + return + if payment.status not in {PaymentModel.SUCCEEDED, PaymentModel.CANCELLED}: + return + PaymentPlanUserInfo.objects.filter(pk=self.user.payment_plan.pk).update( + last_recovery_payment_id=UUID(str(payment.id)), + ) + def do_payment(self, payment: YookassaPaymentResponse) -> PaymentModel: from payments.services.payment_plan_service import PaymentPlanService @@ -65,6 +175,7 @@ class PaymentService: payment_instance, should_process = self.save_payment(payment) if not should_process: return payment_instance + self._close_recovery_generation(payment) logger.info( 'Processing payment webhook: payment_id=%s email=%s status=%s', payment.id, @@ -0,0 +1,407 @@ +from datetime import timedelta +from types import SimpleNamespace +from unittest.mock import patch +from uuid import UUID, uuid4 + +from django.utils import timezone + +from core import tests as core_tests +from payments.models import Payment, PaymentMethod, PaymentPlan, PaymentPlanUserInfo +from payments.services.payment_service import PaymentService + + +class RestoreSubscriptionAPITest(core_tests.BaseAuthorizedAPITest): + ENDPOINT = '/api/v1/payments/restore-subscription' + + @classmethod + def setUpTestData(cls) -> None: + PaymentPlan.objects.update_or_create(price=0, tokens_per_plan=10, defaults={}) + super().setUpTestData() + + @classmethod + def setup_test_data(cls) -> None: + cls.plan = PaymentPlan.objects.create(price=1000, tokens_per_plan=100) + cls.user.payment_plan.plan = cls.plan + cls.user.payment_plan.save() + cls.payment_method = PaymentMethod.objects.create( + user_plan_info=cls.user.payment_plan, + gateway=PaymentMethod.GatewayChoices.BANK_CARD, + payment_method_id=uuid4(), + metadata={}, + active=True, + primary=True, + ) + + def _payment(self, status: str, payment_id: str | None = None) -> SimpleNamespace: + return SimpleNamespace( + amount=SimpleNamespace(value=self.plan.price), + description=str(self.user.uid), + id=payment_id or str(uuid4()), + metadata={'plan_uid': str(self.plan.uid), 'recovery': True}, + status=status, + ) + + def _create_local_payment(self, payment_id: str) -> Payment: + return Payment.objects.create( + uid=payment_id, + user=self.user, + amount=self.plan.price, + plan=self.plan, + status=Payment.SUCCEEDED, + description=str(self.user.uid), + ) + + def _apply_recovery_webhook(self, payment: SimpleNamespace) -> None: + PaymentService(self.user)._close_recovery_generation(payment) + + def test_unauthorized_status_code(self) -> None: + response = self.client.post(self.ENDPOINT) + + self.assertEqual(response.status_code, 401) + self.assertEqual(response.json(), {'detail': 'Unauthorized'}) + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_authorized_status_code(self, create_payment_mock, sleep_mock) -> None: + yookassa_payment = self._payment('succeeded') + create_payment_mock.return_value = yookassa_payment + self._create_local_payment(yookassa_payment.id) + + response = self.post() + + self.assertEqual(response.status_code, 200) + self.assertEqual(response.json(), {'ok': True, 'external': False, 'internal': False}) + sleep_mock.assert_not_called() + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_uses_primary_active_payment_method_first(self, create_payment_mock, sleep_mock) -> None: + newer_payment_method = PaymentMethod( + user_plan_info=self.user.payment_plan, + gateway=PaymentMethod.GatewayChoices.BANK_CARD, + payment_method_id=uuid4(), + metadata={}, + active=True, + primary=False, + ) + PaymentMethod.objects.bulk_create([newer_payment_method]) + yookassa_payment = self._payment('succeeded') + create_payment_mock.return_value = yookassa_payment + self._create_local_payment(yookassa_payment.id) + + self.post() + + payment_data = create_payment_mock.call_args.args[0] + self.assertEqual(payment_data['payment_method_id'], self.payment_method.payment_method_id) + self.assertEqual(payment_data['metadata']['recovery'], True) + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_uses_most_recent_active_payment_method(self, create_payment_mock, sleep_mock) -> None: + PaymentMethod.objects.filter(pk=self.payment_method.pk).update(primary=False) + newer_payment_method = PaymentMethod( + user_plan_info=self.user.payment_plan, + gateway=PaymentMethod.GatewayChoices.BANK_CARD, + payment_method_id=uuid4(), + metadata={}, + active=True, + primary=False, + ) + PaymentMethod.objects.bulk_create([newer_payment_method]) + yookassa_payment = self._payment('succeeded') + create_payment_mock.return_value = yookassa_payment + self._create_local_payment(yookassa_payment.id) + + self.post() + + payment_data = create_payment_mock.call_args.args[0] + self.assertEqual(payment_data['payment_method_id'], newer_payment_method.payment_method_id) + + def test_returns_bad_request_without_active_payment_method(self) -> None: + PaymentMethod.objects.update(active=False) + + response = self.post() + + self.assertEqual(response.status_code, 400) + self.assertEqual(response.json(), {'detail': 'Active payment method not found'}) + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_returns_not_ok_when_payment_is_canceled(self, create_payment_mock, sleep_mock) -> None: + create_payment_mock.return_value = self._payment('canceled') + + response = self.post() + + self.assertEqual(response.status_code, 200) + self.assertEqual(response.json(), {'ok': False, 'external': False, 'internal': False}) + sleep_mock.assert_not_called() + self.user.payment_plan.refresh_from_db() + self.assertIsNotNone(self.user.payment_plan.last_recovery_payment_id) + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.find_one') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_returns_not_ok_when_payment_stays_pending( + self, + create_payment_mock, + find_one_mock, + sleep_mock, + ) -> None: + pending_payment = self._payment('pending') + create_payment_mock.return_value = pending_payment + find_one_mock.return_value = pending_payment + + response = self.post() + + self.assertEqual(response.status_code, 200) + self.assertEqual(response.json(), {'ok': False, 'external': True, 'internal': False}) + self.assertEqual(sleep_mock.call_count, 4) + self.assertEqual(find_one_mock.call_count, 4) + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.find_one') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_returns_ok_when_pending_becomes_succeeded_and_local_payment_exists( + self, + create_payment_mock, + find_one_mock, + sleep_mock, + ) -> None: + pending_payment = self._payment('pending') + succeeded_payment = self._payment('succeeded', payment_id=pending_payment.id) + create_payment_mock.return_value = pending_payment + find_one_mock.return_value = succeeded_payment + self._create_local_payment(pending_payment.id) + + response = self.post() + + self.assertEqual(response.status_code, 200) + self.assertEqual(response.json(), {'ok': True, 'external': False, 'internal': False}) + sleep_mock.assert_called_once_with(0.25) + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_returns_internal_true_when_local_payment_is_missing( + self, + create_payment_mock, + sleep_mock, + ) -> None: + create_payment_mock.return_value = self._payment('succeeded') + + response = self.post() + + self.assertEqual(response.status_code, 200) + self.assertEqual(response.json(), {'ok': False, 'external': False, 'internal': True}) + self.assertEqual(sleep_mock.call_count, 4) + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_returns_forbidden_when_recovery_locked( + self, + create_payment_mock, + sleep_mock, + ) -> None: + self.user.payment_plan.recovery_locked_at = timezone.now() + self.user.payment_plan.save(update_fields=['recovery_locked_at']) + + response = self.post() + + self.assertEqual(response.status_code, 403) + self.assertEqual(response.json(), {'detail': 'Subscription recovery is already in progress'}) + create_payment_mock.assert_not_called() + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_releases_lock_after_request( + self, + create_payment_mock, + sleep_mock, + ) -> None: + yookassa_payment = self._payment('succeeded') + create_payment_mock.return_value = yookassa_payment + self._create_local_payment(yookassa_payment.id) + + response = self.post() + + self.assertEqual(response.status_code, 200) + self.user.payment_plan.refresh_from_db() + self.assertIsNone(self.user.payment_plan.recovery_locked_at) + self.assertIsNone(self.user.payment_plan.last_recovery_payment_id) + + self._apply_recovery_webhook(yookassa_payment) + self.user.payment_plan.refresh_from_db() + self.assertEqual(self.user.payment_plan.last_recovery_payment_id, UUID(str(yookassa_payment.id))) + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_allows_request_when_recovery_lock_is_stale( + self, + create_payment_mock, + sleep_mock, + ) -> None: + yookassa_payment = self._payment('succeeded') + create_payment_mock.return_value = yookassa_payment + self._create_local_payment(yookassa_payment.id) + self.user.payment_plan.recovery_locked_at = timezone.now() - timedelta(minutes=2) + self.user.payment_plan.save(update_fields=['recovery_locked_at']) + + response = self.post() + + self.assertEqual(response.status_code, 200) + self.assertEqual(response.json(), {'ok': True, 'external': False, 'internal': False}) + self.user.payment_plan.refresh_from_db() + self.assertIsNone(self.user.payment_plan.recovery_locked_at) + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_retries_same_idempotency_key_when_local_is_missing( + self, + create_payment_mock, + sleep_mock, + ) -> None: + yookassa_payment = self._payment('succeeded') + create_payment_mock.return_value = yookassa_payment + + first_response = self.post() + self.user.payment_plan.refresh_from_db() + self.assertIsNone(self.user.payment_plan.last_recovery_payment_id) + self.assertIsNone(self.user.payment_plan.recovery_locked_at) + + self._create_local_payment(yookassa_payment.id) + second_response = self.post() + + self.assertEqual(first_response.status_code, 200) + self.assertEqual(first_response.json(), {'ok': False, 'external': False, 'internal': True}) + self.assertEqual(second_response.status_code, 200) + self.assertEqual(second_response.json(), {'ok': True, 'external': False, 'internal': False}) + self.assertEqual(create_payment_mock.call_count, 2) + self.assertEqual( + create_payment_mock.call_args_list[0].kwargs['idempotency_key'], + create_payment_mock.call_args_list[1].kwargs['idempotency_key'], + ) + self.user.payment_plan.refresh_from_db() + self.assertIsNone(self.user.payment_plan.last_recovery_payment_id) + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_webhook_closes_generation_after_internal_true( + self, + create_payment_mock, + sleep_mock, + ) -> None: + first_payment = self._payment('succeeded') + second_payment = self._payment('succeeded') + create_payment_mock.side_effect = [first_payment, second_payment] + + first_response = self.post() + self.assertEqual(first_response.json(), {'ok': False, 'external': False, 'internal': True}) + + self._create_local_payment(first_payment.id) + self._apply_recovery_webhook(first_payment) + self.user.payment_plan.refresh_from_db() + self.assertEqual(self.user.payment_plan.last_recovery_payment_id, UUID(str(first_payment.id))) + + self._create_local_payment(second_payment.id) + second_response = self.post() + + self.assertEqual(second_response.status_code, 200) + self.assertEqual(second_response.json(), {'ok': True, 'external': False, 'internal': False}) + self.assertEqual(create_payment_mock.call_count, 2) + self.assertNotEqual( + create_payment_mock.call_args_list[0].kwargs['idempotency_key'], + create_payment_mock.call_args_list[1].kwargs['idempotency_key'], + ) + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_duplicate_webhook_skips_recovery_close( + self, + create_payment_mock, + sleep_mock, + ) -> None: + yookassa_payment = self._payment('succeeded') + create_payment_mock.return_value = yookassa_payment + self._create_local_payment(yookassa_payment.id) + + self.post() + self._apply_recovery_webhook(yookassa_payment) + self.user.payment_plan.refresh_from_db() + self.assertEqual(self.user.payment_plan.last_recovery_payment_id, UUID(str(yookassa_payment.id))) + + PaymentPlanUserInfo.objects.filter(pk=self.user.payment_plan.pk).update( + last_recovery_payment_id=None + ) + self.user.payment_plan.refresh_from_db() + self.assertIsNone(self.user.payment_plan.last_recovery_payment_id) + + PaymentService(self.user).do_payment(yookassa_payment) + self.user.payment_plan.refresh_from_db() + self.assertIsNone(self.user.payment_plan.last_recovery_payment_id) + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_creates_new_payment_when_previous_recovery_already_applied( + self, + create_payment_mock, + sleep_mock, + ) -> None: + first_payment = self._payment('succeeded') + second_payment = self._payment('succeeded') + create_payment_mock.side_effect = [first_payment, second_payment] + self._create_local_payment(first_payment.id) + + first_response = self.post() + self._apply_recovery_webhook(first_payment) + self._create_local_payment(second_payment.id) + second_response = self.post() + + self.assertEqual(first_response.status_code, 200) + self.assertEqual(first_response.json(), {'ok': True, 'external': False, 'internal': False}) + self.assertEqual(second_response.status_code, 200) + self.assertEqual(second_response.json(), {'ok': True, 'external': False, 'internal': False}) + self.assertEqual(create_payment_mock.call_count, 2) + self.assertNotEqual( + create_payment_mock.call_args_list[0].kwargs['idempotency_key'], + create_payment_mock.call_args_list[1].kwargs['idempotency_key'], + ) + + @patch('payments.services.payment_service.time.sleep') + @patch('payments.services.payment_service.YookassaPayment.create') + def test_creates_new_payment_after_canceled_recovery( + self, + create_payment_mock, + sleep_mock, + ) -> None: + canceled_payment = self._payment('canceled') + succeeded_payment = self._payment('succeeded') + create_payment_mock.side_effect = [canceled_payment, succeeded_payment] + self._create_local_payment(succeeded_payment.id) + + first_response = self.post() + self.user.payment_plan.refresh_from_db() + self.assertEqual(self.user.payment_plan.last_recovery_payment_id, UUID(str(canceled_payment.id))) + self.assertIsNone(self.user.payment_plan.recovery_locked_at) + + second_response = self.post() + + self.assertEqual(first_response.status_code, 200) + self.assertEqual(first_response.json(), {'ok': False, 'external': False, 'internal': False}) + self.assertEqual(second_response.status_code, 200) + self.assertEqual(second_response.json(), {'ok': True, 'external': False, 'internal': False}) + self.assertEqual(create_payment_mock.call_count, 2) + self.assertNotEqual( + create_payment_mock.call_args_list[0].kwargs['idempotency_key'], + create_payment_mock.call_args_list[1].kwargs['idempotency_key'], + ) + + def test_restore_subscription_blocked_endpoint(self) -> None: + response = self.get(endpoint='/api/v1/payments/restore-subscription/blocked') + self.assertEqual(response.status_code, 200) + self.assertEqual(response.json(), {'blocked': False}) + + self.user.payment_plan.recovery_locked_at = timezone.now() + self.user.payment_plan.save(update_fields=['recovery_locked_at']) + + response = self.get(endpoint='/api/v1/payments/restore-subscription/blocked') + self.assertEqual(response.status_code, 200) + self.assertEqual(response.json(), {'blocked': True}) @@ -49,6 +49,16 @@ class PaymentLinkSchema(Schema): payment_url: str +class RestoreSubscriptionSchema(Schema): + ok: bool + external: bool + internal: bool + + +class RestoreSubscriptionLockSchema(Schema): + blocked: bool + + class UserPlanDetailSchema(Schema): uid: UUID plan: PaymentPlanSchema @@ -84,5 +94,3 @@ class ExpensesParamsSchema(Schema): class ExpensesSchema(Schema): source: str amount: condecimal(max_digits=10, decimal_places=2) - - @@ -1,3 +1,4 @@ +import logging from decimal import Decimal from celery import shared_task @@ -14,6 +15,8 @@ from tools.chats.services.sse_chunk_service import SSEChunkService from tools.chats.services.sse_store import PublicSSEStoreService, SSEStoreService from tools.public_api.models import APIKey, APIStore +logger = logging.getLogger(__name__) + def _run_stream( store: SSEStoreService, @@ -51,6 +54,7 @@ def _run_stream( event_id += 1 store.push(SSEChunkService.done(event_id, exc.value or ''), ttl=settings.SSE_DONE_STREAM_TTL) except Exception as exc: + logger.exception(f'Model streaming failed: {(exc.__cause__ or exc)!r}') if event_id < 2: message.is_sent = False message.save(update_fields=['is_sent'])