@@ -9,6 +9,7 @@ from backend import settings from ml_model.exceptions import ( FileExtensionNotSupported, GenerationException, + NSFWDetectedException, RealPersonDetectedError, RequestBlocked, ) @@ -284,6 +285,10 @@ class BytedanceModelArkAdapter: resp.text, ) raise GenerationException from exc + if error_code := data.get('error', {}).get('code', ''): + if error_code == 'OutputImageSensitiveContentDetected': + raise NSFWDetectedException + if image_data := data.get('data'): urls = [item.get('url') for item in image_data if isinstance(item, dict) and item.get('url')] if urls: @@ -0,0 +1,8 @@ +from dataclasses import dataclass + + +@dataclass +class ModelResponse: + content: str + input_tokens: int + output_tokens: int @@ -0,0 +1,121 @@ +import json +import logging +import math +from typing import Any + +import httpx +import tiktoken + +from backend import settings +from poller.models import Proxy + +from .models import ModelResponse + +logger = logging.getLogger(__name__) + + +class OpenrouterAdapter: + BASE_URL = 'https://openrouter.ai/api/v1' + FALLBACK_ENCODINGS: dict[str, str] = { + 'Qwen': 'cl100k_base', + 'Deepseek': 'cl100k_base', + 'Claude': 'r50k_base', + 'Perplexity': 'cl100k_base', + 'Mistral': 'r50k_base', + 'LLaMA': 'cl100k_base', + 'Grok': 'r50k_base', + 'Gemini': 'cl100k_base', + } + FALLBACK_TOKEN_BIAS: dict[str, tuple[float, float]] = { + # Средний калибровочный bias по 15 замерам: + # input: 2814 / 3246 ~= 0.867, output: 14223 / 25767 ~= 0.552 + 'Grok': (0.867, 0.552), + } + + @classmethod + def run_streaming_api( + cls, version: str, messages: list, callback_data: dict, model_name: str + ) -> ModelResponse: + for proxy in Proxy.objects.all(): + with httpx.Client( + base_url=cls.BASE_URL, + headers={'Authorization': f'Bearer {settings.OPENROUTER_API_KEY}'}, + proxy=f'{proxy.protocol}://{proxy.address}', + timeout=600, + ) as client: + with client.stream( + 'POST', + 'chat/completions', + json={ + 'model': version, + 'stream': True, + 'messages': messages, + 'transforms': ['middle-out'], + **callback_data, + }, + ) as resp: + content = '' + # reasoning используем только для фоллбэк-подсчёта токенизатора + # в ответ не кладём, заполняет буфер истории сообщений + reasoning = '' + input_tokens = output_tokens = 0 + for line in resp.iter_lines(): + line = line.strip() + if not line or not line.startswith('data: '): + continue + + data = line[6:] + if data == '[DONE]': + break + + try: + data_obj = json.loads(data) + content += data_obj['choices'][0]['delta'].get('content') or '' + reasoning += data_obj['choices'][0]['delta'].get('reasoning') or '' + if data_obj.get('usage'): + input_tokens = data_obj['usage']['prompt_tokens'] + output_tokens = data_obj['usage']['completion_tokens'] + + except json.JSONDecodeError: + logger.warning( + f'Opernrouter chunk parsing failed for model {model_name}: {data}' + ) + continue + + if not input_tokens or not output_tokens: + logger.error(f'Opernrouter failed get data about tokens for model {model_name}') + input_tokens, output_tokens = cls._fallback_tokenize( + model_name, messages, content + reasoning + ) + + return ModelResponse(content, input_tokens, output_tokens) + + @classmethod + def _fallback_tokenize(cls, model_name: str, messages: list, content: str) -> tuple[int, int]: + bias = cls.FALLBACK_TOKEN_BIAS.get(model_name, (1.0, 1.0)) + return cls.count_tokens_fallback(model_name, messages, content, bias=bias) + + @classmethod + def count_tokens_fallback( + cls, + model_name: str, + messages: list[dict[str, Any]], + output: str, + bias: tuple[float, float] = (1.0, 1.0), + ) -> tuple[int, int]: + """Fallback token counting for OpenRouter responses.""" + encoding = tiktoken.get_encoding(cls.FALLBACK_ENCODINGS[model_name]) + input_tokens = 100 if model_name == 'LLaMA' else 0 + for message in messages: + content = message['content'] + if isinstance(content, list): + input_tokens += len(encoding.encode(content[0]['text'])) + else: + input_tokens += len(encoding.encode(content)) + output_tokens = len(encoding.encode(output)) + + input_bias, output_bias = bias + input_tokens = math.ceil(input_tokens * input_bias) + output_tokens = math.ceil(output_tokens * output_bias) + + return (input_tokens, output_tokens) @@ -3,23 +3,21 @@ import time from datetime import timedelta from decimal import Decimal from io import BytesIO -from typing import Any, Iterator from pathlib import Path +from typing import Any, Iterator import filetype from PIL import Image -from ml_model.services.FileService import FileProcessingService + +from messages.models import Message from ml_model.exceptions import CorruptedFileError, FileExtensionNotSupported +from ml_model.services.base import SimpleService from ml_model.services.EmbeddingService import EmbeddingService +from ml_model.services.FileService import FileProcessingService +from ml_model.tasks import openrouter_run_streaming from payments.exceptions.insufficient_balance import InsufficientBalance from payments.selectors.payment_plan_selector import PaymentPlanSelector - - from poller.models import Proxy - -from messages.models import Message -from ml_model.services.base import SimpleService -from ml_model.tasks import openrouter_run from tools.chats.models import Chat from tools.copywrite.models import Copywrite from tools.public_api.models import APIStore @@ -43,9 +41,7 @@ class Grok(SimpleService): ) -> Decimal: price_map = self.TOKENS_COST[version.split('/')[1]] price = ( - input_tokens - * price_map['input']['default' if input_tokens <= 200_000 else 'high'] - / 1_000_000 + input_tokens * price_map['input']['default' if input_tokens <= 200_000 else 'high'] / 1_000_000 + output_tokens * price_map['output']['default' if output_tokens <= 200_000 else 'high'] / 1_000_000 @@ -134,16 +130,16 @@ class Grok(SimpleService): else: raise FileExtensionNotSupported(self.SUPPORTED_EXTENSIONS) - result = openrouter_run(version, messages, callback_data, 'Grok') + result = openrouter_run_streaming(version, messages, callback_data, 'Grok') process_time = timedelta(seconds=(time.time() - start_time)) self.handle_invoice( input_message.content_object.model, version=version, - input_tokens=result[1], - output_tokens=result[2], + input_tokens=result.input_tokens, + output_tokens=result.output_tokens, embedding_tokens=embedding_tokens, ) - msgs = self.save_results(result[0], process_time) + msgs = self.save_results(result.content, process_time) return msgs def get_chat_history( @@ -0,0 +1 @@ + @@ -1,5 +1,6 @@ from typing import Type +from django.db import transaction from django.db.models.signals import post_save from django.dispatch import receiver @@ -9,4 +10,4 @@ from ml_model.models import ModelSettings, NeuronModel @receiver(post_save, sender=NeuronModel) def create_settings(sender: Type[NeuronModel], instance: NeuronModel, created: bool, **kwargs): if created: - ModelSettings.objects.create(model=instance) + transaction.on_commit(lambda: ModelSettings.objects.get_or_create(model=instance)) @@ -21,6 +21,7 @@ from requests import Response from backend import settings from ml_model.adapters.bytedance_model_ark import BytedanceContentType, BytedanceModelArkAdapter +from ml_model.adapters.openrouter import OpenrouterAdapter from ml_model.exceptions import ( CorruptedFileError, DeploymentDisabled, @@ -35,7 +36,6 @@ from ml_model.exceptions import ( PredictionInterruptedError, RequestBlocked, ) -from ml_model.utils import count_openrouter_tokens from poller.models import Proxy logger = logging.getLogger(__name__) @@ -189,7 +189,7 @@ def openrouter_run(version: str, messages: list, callback_data: dict, model_name error_type = re.sub(r'["\']', '', str(data['choices'][0]['error']['message'])) if error_type == 'Overloaded': logger.warning(f'Model {model_name} overloaded') - input_tokens, output_tokens = count_openrouter_tokens( + input_tokens, output_tokens = OpenrouterAdapter.count_tokens_fallback( model_name, messages, content + reasoning ) else: @@ -203,6 +203,11 @@ def openrouter_run(version: str, messages: list, callback_data: dict, model_name raise Exception(f'No answer from {model_name}, please retry later') +@shared_task +def openrouter_run_streaming(version: str, messages: list, callback_data: dict, model_name: str): + return OpenrouterAdapter.run_streaming_api(version, messages, callback_data, model_name) + + @shared_task def fal_ai_run(model, payload): for proxy in Proxy.objects.all(): @@ -1,7 +1,5 @@ from random import randint -from typing import Any, Dict, List, Literal, Tuple - -import tiktoken +from typing import Literal from authentication.models import CustomUserModel from authentication.selectors.account_status_selector import ( @@ -32,26 +30,3 @@ def check_account_type( return 'business_admin' return 'business_account' - - -def count_openrouter_tokens(model_name: str, messages: List[Dict[str, Any]], output: str) -> Tuple[int, int]: - """A function for count tokens for OpenRouter Neuron Models""" - encodings = { - 'Qwen': 'cl100k_base', - 'Deepseek': 'cl100k_base', - 'Claude': 'r50k_base', - 'Perplexity': 'cl100k_base', - 'Mistral': 'r50k_base', - 'LLaMA': 'cl100k_base', - 'Grok': 'r50k_base', - 'Gemini': 'cl100k_base', - } - encoding = tiktoken.get_encoding(encodings[model_name]) - input_tokens = 100 if model_name == 'LLaMA' else 0 - for message in messages: - if isinstance(message['content'], list): - input_tokens += len(encoding.encode(message['content'][0]['text'])) - else: - input_tokens += len(encoding.encode(message['content'])) - output_tokens = len(encoding.encode(output)) - return (input_tokens, output_tokens) @@ -0,0 +1,250 @@ +import argparse +import difflib +import re +from datetime import datetime +from decimal import Decimal, InvalidOperation +from pathlib import Path +from typing import Any, NamedTuple + +from django.core.management.base import BaseCommand, CommandError +from django.db.migrations.loader import MigrationLoader + +from ml_model.models import NeuronModel +from payments.models import PaymentPlanFeature + +MEASUREMENT_UNITS = {choice.value for choice in PaymentPlanFeature.MeasurementUnitChoices} +PAYMENTS_APP_LABEL = 'payments' + + +def _get_model_slugs() -> set[str]: + return set(NeuronModel.objects.values_list('slug', flat=True)) + + +def _parse_model_slug(value: str, model_slugs: set[str]) -> str: + model_slug = value.strip() + if model_slug in model_slugs: + return model_slug + close = difflib.get_close_matches(model_slug, model_slugs, n=1) + if close: + raise argparse.ArgumentTypeError( + f'NeuronModel со slug "{model_slug}" не найдена. ' + f'Возможно, вы имели в виду: {close[0]}' + ) + raise argparse.ArgumentTypeError(f'NeuronModel со slug "{model_slug}" не найдена.') + + +def _sanitize_slug(slug: str) -> str: + return re.sub(r'[^a-zA-Z0-9_]', '_', slug) + + +class PaymentsMigrationSlot(NamedTuple): + """dependency — лист графа для dependencies; next_number — префикс нового файла.""" + + dependency: str + next_number: int + + +def _resolve_payments_migration_slot() -> PaymentsMigrationSlot: + """Один проход MigrationLoader: хвост ветки и номер следующей миграции.""" + loader = MigrationLoader(None, ignore_no_migrations=True) + + leaves = sorted(name for app, name in loader.graph.leaf_nodes() if app == PAYMENTS_APP_LABEL) + if not leaves: + raise CommandError(f'В приложении {PAYMENTS_APP_LABEL} нет миграций.') + if len(leaves) > 1: + raise CommandError( + f'У {PAYMENTS_APP_LABEL} несколько концов веток: {", ".join(leaves)}. ' + f'Сначала смержите: python manage.py makemigrations {PAYMENTS_APP_LABEL} --merge' + ) + + dependency = leaves[0] + prefix = dependency.split('_', 1)[0] + if not prefix.isdigit(): + raise CommandError(f'Не удалось вычислить номер следующей миграции из листа "{dependency}".') + return PaymentsMigrationSlot(dependency=dependency, next_number=int(prefix) + 1) + + +def _migration_basename(model_slug: str, next_number: int) -> str: + safe_slug = _sanitize_slug(model_slug) + return f'{next_number:04d}_add_{safe_slug}_payment_features' + + +def _build_migration_source( + *, + func_name: str, + model_slug: str, + price: Decimal, + measurement_unit: str, + price_threshold: int, + latest_migration: str, +) -> str: + return f"""# Generated by makemigration_payment_features on {datetime.now():%Y-%m-%d %H:%M} + +import math +from decimal import Decimal + +from django.db import migrations +from django.db.models import Max + + +def {func_name}(apps, schema_editor): + PaymentPlan = apps.get_model('payments', 'PaymentPlan') + PaymentPlanFeature = apps.get_model('payments', 'PaymentPlanFeature') + NeuronModel = apps.get_model('ml_model', 'NeuronModel') + + price = Decimal('{price}') + measurement_unit = '{measurement_unit}' + price_threshold = {price_threshold} + model = NeuronModel.objects.get(slug='{model_slug}') + category = model.category + + max_order_by_plan_id = {{ + row['plan_id']: row['max_order'] + for row in PaymentPlanFeature.objects.filter(model__category=category) + .values('plan_id') + .annotate(max_order=Max('order')) + }} + + features = [] + for plan in PaymentPlan.objects.filter(price__gt=price_threshold): + quantity = math.floor(plan.tokens_per_plan / price) + max_order = max_order_by_plan_id.get(plan.pk) + next_order = (max_order if max_order is not None else -1) + 1 + max_order_by_plan_id[plan.pk] = next_order + features.append( + PaymentPlanFeature( + plan=plan, + model=model, + quantity=quantity, + measurement_unit=measurement_unit, + order=next_order, + ) + ) + PaymentPlanFeature.objects.bulk_create( + features, + update_conflicts=True, + update_fields=['quantity', 'measurement_unit'], + unique_fields=['plan', 'model'], + ) + + +class Migration(migrations.Migration): + + dependencies = [ + ('payments', '{latest_migration}'), + ] + + operations = [ + migrations.RunPython({func_name}, migrations.RunPython.noop), + ] +""" + + +class Command(BaseCommand): + help = ( + 'Создаёт data-миграцию payments с PaymentPlanFeature для указанной модели. ' + 'Применение: python manage.py migrate payments' + ) + + def add_arguments(self, parser): + model_slugs = _get_model_slugs() + parser.add_argument( + 'modelname', + type=lambda value: _parse_model_slug(value, model_slugs), + metavar='model_slug', + help=f'Slug модели (NeuronModel). Доступные: {", ".join(sorted(model_slugs))}', + ) + parser.add_argument('--price', type=str, help='Цена за единицу измерения (токены)') + parser.add_argument( + '--measurement-unit', + type=str, + choices=sorted(MEASUREMENT_UNITS), + help='Единица измерения: text_page, file, time', + ) + parser.add_argument( + '--price-threshold', + type=int, + help='Создавать фичи только для планов с price строго больше этого значения', + ) + + def handle(self, *args: Any, **options: Any) -> None: + model_slug = options['modelname'] + price = self._read_price(options.get('price')) + measurement_unit = self._read_measurement_unit(options.get('measurement_unit')) + price_threshold = self._read_price_threshold(options.get('price_threshold')) + + self.stdout.write( + f'Миграция для model={model_slug}, price={price}, ' + f'measurement_unit={measurement_unit}, price_threshold={price_threshold}' + ) + self.stdout.write('Расчёт quantity выполнится на стейдже/проде при migrate по планам в БД.') + + slot = _resolve_payments_migration_slot() + migration_name = _migration_basename(model_slug, slot.next_number) + func_name = f'add_{_sanitize_slug(model_slug)}_payment_features' + source = _build_migration_source( + func_name=func_name, + model_slug=model_slug, + price=price, + measurement_unit=measurement_unit, + price_threshold=price_threshold, + latest_migration=slot.dependency, + ) + + migration_path = ( + Path(__file__).resolve().parent.parent.parent / 'migrations' / f'{migration_name}.py' + ) + if migration_path.exists(): + raise CommandError(f'Файл миграции уже существует: {migration_path}') + + migration_path.write_text(source, encoding='utf-8') + self.stdout.write(self.style.SUCCESS(f'Создана миграция: {migration_path}')) + self.stdout.write('Примените: python manage.py migrate payments') + + def _read_price(self, value: str | None) -> Decimal: + if value is not None: + try: + price = Decimal(value) + except InvalidOperation as exc: + raise CommandError(f'Некорректная цена: {value}') from exc + if price <= 0: + raise CommandError('Цена должна быть больше 0.') + return price + + while True: + raw = input('Введите прайс за единицу измерения (токены): ').strip() + try: + price = Decimal(raw) + except InvalidOperation: + self.stderr.write('Введите число.') + continue + if price <= 0: + self.stderr.write('Цена должна быть больше 0.') + continue + return price + + def _read_measurement_unit(self, value: str | None) -> str: + units = ', '.join(sorted(MEASUREMENT_UNITS)) + + if value is not None: + if value not in MEASUREMENT_UNITS: + raise CommandError(f'Недопустимая ед. измерения: {value}. Допустимые: {units}') + return value + + while True: + raw = input(f'Введите ед. измерения ({units}): ').strip() + if raw in MEASUREMENT_UNITS: + return raw + self.stderr.write(f'Допустимые значения: {units}') + + def _read_price_threshold(self, value: int | None) -> int: + if value is not None: + return value + + while True: + raw = input('Введите планы выше какого price будут учитываться: ').strip() + try: + return int(raw) + except ValueError: + self.stderr.write('Введите целое число.') + @@ -25,9 +25,7 @@ from payments.models import ( PaymentPlan, PaymentPlanFeature, PaymentMethod, - PaymentPlanUserInfo, ) -from authentication.services.email_service import EmailService from payments.schema import UserBalance from payments.schemas import ( ExpensesParamsSchema, @@ -228,25 +226,28 @@ async def create_payment_link(request, body: NewSubscriptionSchema): @router.post('gitlab-webhook', tags=['payments/gitlab-webhook'], auth=None) async def handle_gitlab_webhook(request): - try: - data = orjson.loads(request.body)['object_attributes'] - if data['name'] == 'recurring_payments': - is_active = data['active'] - if not is_active: - deleted_methods_count, deleted_details = await PaymentMethod.objects.all().adelete() - logger.info( - 'Recurring feature disabled: all payment methods removed count=%s', - deleted_methods_count, - ) - updated_count = await PaymentPlanUserInfo.objects.filter( - plan__price__gt=0, - plan__individual=False, - ).aupdate(next_payment_at=None if not is_active else (timezone.now() + timedelta(days=30))) - logger.info( - 'Recurring feature flag synced: active=%s updated_subscriptions=%s', - is_active, - updated_count, - ) - except Exception as exc: - logger.error(exc) + ''' + Currently disabled, pending future feature flags + ''' + # try: + # data = orjson.loads(request.body)['object_attributes'] + # if data['name'] == 'recurring_payments': + # is_active = data['active'] + # if not is_active: + # deleted_methods_count, deleted_details = await PaymentMethod.objects.all().adelete() + # logger.info( + # 'Recurring feature disabled: all payment methods removed count=%s', + # deleted_methods_count, + # ) + # updated_count = await PaymentPlanUserInfo.objects.filter( + # plan__price__gt=0, + # plan__individual=False, + # ).aupdate(next_payment_at=None if not is_active else (timezone.now() + timedelta(days=30))) + # logger.info( + # 'Recurring feature flag synced: active=%s updated_subscriptions=%s', + # is_active, + # updated_count, + # ) + # except Exception as exc: + # logger.error(exc) return 200 @@ -14,7 +14,6 @@ from yookassa import Payment as YookassaPayment from yookassa.domain.response import PaymentResponse as YookassaPaymentResponse from authentication.models import CustomUserModel -from lib.unleash.client import web_client from payments.models.payment import Payment as PaymentModel from payments.models.payment_plan import PaymentPlan, PaymentPlanUserInfo from payments.services.payment_method_service import PaymentMethodService @@ -42,11 +41,6 @@ class PaymentService: } ], } - is_recurring = ( - web_client.get_flag_state('recurring_payments', self.user.email) - and not plan.individual - and web_client.get_flag_state('auto-save-payments-enabled', self.user.email) - ) payment_data = { 'amount': {'value': f'{plan.price}', 'currency': 'RUB'}, 'receipt': receipt_data, @@ -56,16 +50,12 @@ class PaymentService: }, 'description': str(self.user.uid), 'capture': True, - 'save_payment_method': is_recurring, + 'save_payment_method': True, 'metadata': {'plan_uid': str(plan.uid)}, } payment = YookassaPayment.create(payment_data, uuid4()) logger.info( - 'Payment link created: email=%s plan_uid=%s price=%s recurring=%s', - self.user.email, - plan.uid, - plan.price, - is_recurring, + 'Payment link created: email=%s plan_uid=%s price=%s', self.user.email, plan.uid, plan.price ) return payment.confirmation.confirmation_url @@ -113,11 +103,7 @@ class PaymentService: return self.user.payment_plan.current_token_balance + plan.tokens_per_plan def handle_succeeded_payment(self, payment: YookassaPaymentResponse, plan: PaymentPlan) -> None: - if ( - payment.payment_method.saved - and web_client.get_flag_state('recurring_payments', self.user.email) - and not plan.individual - ): + if payment.payment_method.saved and not plan.individual: payment_method = PaymentMethodService(self.user).add_payment_method(payment.payment_method) PaymentPlanUserInfo.objects.filter(user=self.user).update( method=payment_method, next_payment_at=timezone.now() + timedelta(days=30) @@ -128,7 +114,7 @@ class PaymentService: payment_method.uid, ) else: - if web_client.get_flag_state('recurring_payments', self.user.email) and not plan.individual: + if not plan.individual: PaymentPlanUserInfo.objects.filter(user=self.user).update( next_payment_at=timezone.now() + timedelta(days=30) ) @@ -139,7 +125,7 @@ class PaymentService: else: PaymentPlanUserInfo.objects.filter(user=self.user).update(next_payment_at=None) logger.info( - 'Recurring schedule cleared: email=%s reason=feature_disabled_or_individual_plan', + 'Recurring schedule cleared: email=%s reason=individual_plan', self.user.email, ) PaymentMethodService(self.user).delete_payment_method() @@ -14,8 +14,7 @@ from django.utils import timezone from authentication.models.business_host import BusinessUserHost from authentication.models.user import CustomUserModel from authentication.services.email_service import EmailService -from lib.unleash.client import celery_client -from payments.models import PaymentMethod, PaymentPlan, PaymentPlanUserInfo +from payments.models import PaymentPlan, PaymentPlanUserInfo from payments.services.payment_plan_service import PaymentPlanService from yookassa import Payment as YookassaPayment @@ -44,11 +43,6 @@ def withdraw(user_id: UUID, amount: Decimal): @shared_task def execute_recurring_payments() -> None: - if not celery_client.is_feature_enabled('recurring_payments'): - logger.info('Recurring payments feature disabled, skipping execute') - return - - emails = celery_client.get_user_emails('recurring_payments') overdue_payments = ( PaymentPlanUserInfo.objects.select_related('user', 'plan', 'method') .filter( @@ -75,8 +69,6 @@ def execute_recurring_payments() -> None: 'method__attempts', ) ) - if emails: - overdue_payments = overdue_payments.filter(user__email__in=emails) for overdue_payment in overdue_payments.iterator(chunk_size=CHUNK_SIZE): customer = overdue_payment.user plan = overdue_payment.plan @@ -126,54 +118,26 @@ def revoke_recurring_payments() -> None: logger.info('Revoke recurring already running, skipping') return try: - feature_name = 'recurring_payments' - base_qs = PaymentPlanUserInfo.objects.filter( + qs = PaymentPlanUserInfo.objects.filter( next_payment_at__isnull=False, next_payment_at__lte=timezone.now(), plan__price__gt=0, plan__individual=False, plan__is_corporate=False, + method__isnull=True, ) - revoked_count = 0 - canceled_count = 0 - deleted_methods_count = 0 - if not celery_client.is_feature_enabled(feature_name): - flag_off_qs = base_qs - flag_on_qs = base_qs.none() - else: - allowed_emails = celery_client.get_user_emails(feature_name) - if allowed_emails: - flag_off_qs = base_qs.exclude(user__email__in=allowed_emails) - flag_on_qs = base_qs.filter(user__email__in=allowed_emails, method__isnull=True) - else: - flag_off_qs = base_qs.none() - flag_on_qs = base_qs.filter(method__isnull=True) + canceled_count = 0 - uid_iter = flag_off_qs.values_list('uid', flat=True).iterator(chunk_size=CHUNK_SIZE) - while uids := list(islice(uid_iter, CHUNK_SIZE)): + free_regular_plan = PaymentPlan.objects.get_or_create(price=0, is_corporate=False)[0] + qs_iter = qs.values_list('uid', flat=True).iterator(chunk_size=CHUNK_SIZE) + while uids := list(islice(qs_iter, CHUNK_SIZE)): with transaction.atomic(): - deleted_methods_count += PaymentMethod.objects.filter( - user_plan_info__uid__in=uids - ).delete()[0] - revoked_count += flag_off_qs.filter(uid__in=uids).update(next_payment_at=None) - - if flag_on_qs.exists(): - free_regular_plan = PaymentPlan.objects.get_or_create(price=0, is_corporate=False)[0] - flag_on_iter = flag_on_qs.values_list('uid', flat=True).iterator(chunk_size=CHUNK_SIZE) - while uids := list(islice(flag_on_iter, CHUNK_SIZE)): - with transaction.atomic(): - canceled_count += flag_on_qs.filter(uid__in=uids).update( - next_payment_at=None, - plan_id=free_regular_plan.pk, - current_token_balance=0, - ) - - logger.info( - 'Revoke recurring finished: revoked=%s free_regular=%s deleted_methods=%s', - revoked_count, - canceled_count, - deleted_methods_count, - ) + canceled_count += qs.filter(uid__in=uids).update( + next_payment_at=None, + plan_id=free_regular_plan.pk, + current_token_balance=0, + ) + logger.info('Revoke recurring finished: free_regular=%s', canceled_count) finally: cache.delete(lock_key)