From dc7b8dc72a3a398d6270a0a2b8ce9e2b54cb9af7 Mon Sep 17 00:00:00 2001 From: Fringg Date: Wed, 4 Mar 2026 07:24:15 +0300 Subject: [PATCH] feat: account linking and merge system for cabinet Add OAuth provider linking/unlinking endpoints, merge token service (Redis-backed, 30-min TTL), and atomic account merge executor that transfers OAuth IDs, telegram_id, email, balance, subscriptions, transactions, payments, referral data, and partner status between two user accounts. Unchosen subscription is deleted from RemnaWave with disable as fallback. Includes 39 unit tests covering all merge scenarios. --- app/cabinet/auth/merge_service.py | 101 +++ app/cabinet/routes/__init__.py | 3 + app/cabinet/routes/account_linking.py | 489 ++++++++++++++ app/services/account_merge_service.py | 483 ++++++++++++++ tests/services/test_account_merge_service.py | 643 +++++++++++++++++++ 5 files changed, 1719 insertions(+) create mode 100644 app/cabinet/auth/merge_service.py create mode 100644 app/cabinet/routes/account_linking.py create mode 100644 app/services/account_merge_service.py create mode 100644 tests/services/test_account_merge_service.py diff --git a/app/cabinet/auth/merge_service.py b/app/cabinet/auth/merge_service.py new file mode 100644 index 00000000..63f24d4a --- /dev/null +++ b/app/cabinet/auth/merge_service.py @@ -0,0 +1,101 @@ +"""Temporary merge token management for account linking. + +Stores short-lived tokens in Redis so the user can confirm merging +two cabinet accounts (primary absorbs secondary) via a separate +confirmation endpoint. +""" + +import secrets +from datetime import UTC, datetime +from typing import Any + +import structlog + +from app.utils.cache import cache, cache_key + + +logger = structlog.get_logger(__name__) + +MERGE_TOKEN_TTL_SECONDS = 1800 # 30 minutes +MERGE_TOKEN_PREFIX = 'account_merge' + + +async def create_merge_token( + primary_user_id: int, + secondary_user_id: int, + provider: str, + provider_id: str, +) -> str: + """Generate a merge token and store its payload in Redis. + + The token is a one-time confirmation handle: whoever presents it + within ``MERGE_TOKEN_TTL_SECONDS`` can execute the account merge. + + Returns the raw token string (URL-safe base64, 32 bytes of entropy). + Raises ``RuntimeError`` if Redis write fails. + """ + token = secrets.token_urlsafe(32) + value: dict[str, Any] = { + 'primary_user_id': primary_user_id, + 'secondary_user_id': secondary_user_id, + 'provider': provider, + 'provider_id': provider_id, + 'created_at': datetime.now(UTC).isoformat(), + } + key = cache_key(MERGE_TOKEN_PREFIX, token) + stored = await cache.set(key, value, expire=MERGE_TOKEN_TTL_SECONDS) + if not stored: + logger.error( + 'Failed to store merge token in Redis', + primary_user_id=primary_user_id, + secondary_user_id=secondary_user_id, + provider=provider, + ) + raise RuntimeError('Failed to store merge token') + + logger.info( + 'Merge token created', + primary_user_id=primary_user_id, + secondary_user_id=secondary_user_id, + provider=provider, + provider_id=provider_id, + ) + return token + + +async def get_merge_token_data(token: str) -> dict[str, Any] | None: + """Read merge token payload *without* consuming it. + + Intended for preview / confirmation screens where the user sees + what will happen before they press "Confirm". + + Returns ``None`` when the token is expired, missing, or malformed. + """ + key = cache_key(MERGE_TOKEN_PREFIX, token) + data: Any = await cache.get(key) + if data is None or not isinstance(data, dict): + return None + return data + + +async def consume_merge_token(token: str) -> dict[str, Any] | None: + """Atomically read and delete a merge token (GETDEL). + + This prevents double-merge race conditions: only the first caller + that reaches Redis will get the payload; every subsequent attempt + receives ``None``. + + Returns the stored dict or ``None`` if already consumed / expired. + """ + key = cache_key(MERGE_TOKEN_PREFIX, token) + data: Any = await cache.getdel(key) + if data is None or not isinstance(data, dict): + return None + + logger.info( + 'Merge token consumed', + primary_user_id=data.get('primary_user_id'), + secondary_user_id=data.get('secondary_user_id'), + provider=data.get('provider'), + ) + return data diff --git a/app/cabinet/routes/__init__.py b/app/cabinet/routes/__init__.py index d1a33622..4ab7376f 100644 --- a/app/cabinet/routes/__init__.py +++ b/app/cabinet/routes/__init__.py @@ -2,6 +2,7 @@ from fastapi import APIRouter +from .account_linking import merge_router as merge_router, router as account_linking_router from .admin_apps import router as admin_apps_router from .admin_audit_log import router as admin_audit_log_router from .admin_ban_system import router as admin_ban_system_router @@ -60,6 +61,8 @@ router = APIRouter(prefix='/cabinet', tags=['Cabinet']) # Include all sub-routers router.include_router(auth_router) router.include_router(oauth_router) +router.include_router(account_linking_router) +router.include_router(merge_router) router.include_router(subscription_router) router.include_router(balance_router) router.include_router(referral_router) diff --git a/app/cabinet/routes/account_linking.py b/app/cabinet/routes/account_linking.py new file mode 100644 index 00000000..40163463 --- /dev/null +++ b/app/cabinet/routes/account_linking.py @@ -0,0 +1,489 @@ +"""Account linking and merge routes for cabinet. + +Router 1 (`router`): JWT-protected endpoints for linking/unlinking OAuth providers. +Router 2 (`merge_router`): Public endpoints for merge preview and execution. +""" + +from datetime import UTC, datetime +from typing import Any + +import structlog +from fastapi import APIRouter, Depends, HTTPException, status +from pydantic import BaseModel, Field +from sqlalchemy.ext.asyncio import AsyncSession + +from app.database.crud.user import ( + get_user_by_id, + get_user_by_oauth_provider, + set_user_oauth_provider_id, +) +from app.database.models import User +from app.services.account_merge_service import execute_merge, get_merge_preview + +from ..auth.merge_service import ( + MERGE_TOKEN_TTL_SECONDS, + consume_merge_token, + create_merge_token, + get_merge_token_data, +) +from ..auth.oauth_providers import ( + generate_oauth_state, + get_provider, + validate_oauth_state, +) +from ..dependencies import get_cabinet_db, get_current_cabinet_user +from ..schemas.auth import UserResponse +from .auth import _create_auth_response, _store_refresh_token, _user_to_response + + +logger = structlog.get_logger(__name__) + +# OAuth provider -> User model column name +_OAUTH_PROVIDER_COLUMNS: dict[str, str] = { + 'google': 'google_id', + 'yandex': 'yandex_id', + 'discord': 'discord_id', + 'vk': 'vk_id', +} + +# All known auth providers (including non-OAuth) +_ALL_PROVIDERS: tuple[str, ...] = ('telegram', 'email', 'google', 'yandex', 'discord', 'vk') + + +# --------------------------------------------------------------------------- +# Schemas +# --------------------------------------------------------------------------- + + +class LinkedProvider(BaseModel): + provider: str + linked: bool + identifier: str | None = None + + +class LinkedProvidersResponse(BaseModel): + providers: list[LinkedProvider] + + +class LinkInitResponse(BaseModel): + authorize_url: str + state: str + + +class LinkCallbackRequest(BaseModel): + code: str = Field(..., min_length=1, max_length=2048, description='Authorization code from provider') + state: str = Field(..., min_length=1, max_length=128, description='CSRF state token') + device_id: str | None = Field(None, max_length=256, description='Device ID from VK ID callback') + + +class LinkCallbackResponse(BaseModel): + success: bool + message: str | None = None + merge_required: bool = False + merge_token: str | None = None + + +class UnlinkResponse(BaseModel): + success: bool + + +class MergePreviewUser(BaseModel): + id: int + username: str | None = None + first_name: str | None = None + email: str | None = None + auth_methods: list[str] + balance_kopeks: int = 0 + subscription: dict[str, Any] | None = None + created_at: datetime | None = None + + +class MergePreviewResponse(BaseModel): + primary: MergePreviewUser + secondary: MergePreviewUser + expires_in_seconds: int + + +class MergeRequest(BaseModel): + keep_subscription_from: int = Field(..., description='User ID whose subscription to keep') + + +class MergeResponse(BaseModel): + success: bool + access_token: str | None = None + refresh_token: str | None = None + user: UserResponse | None = None + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _get_provider_identifier(user: User, provider: str) -> str | None: + """Return the identifier (provider_id or email) for a given provider, or None.""" + match provider: + case 'telegram': + return str(user.telegram_id) if user.telegram_id else None + case 'email': + return user.email if user.email and user.password_hash else None + case _: + column = _OAUTH_PROVIDER_COLUMNS.get(provider) + if not column: + return None + value = getattr(user, column, None) + return str(value) if value else None + + +def _count_auth_methods(user: User) -> int: + """Count how many auth methods the user has linked.""" + count = 0 + if user.telegram_id: + count += 1 + if user.email and user.password_hash: + count += 1 + for column in _OAUTH_PROVIDER_COLUMNS.values(): + if getattr(user, column, None): + count += 1 + return count + + +# --------------------------------------------------------------------------- +# Router 1: Account linking (JWT required) +# --------------------------------------------------------------------------- + +router = APIRouter(prefix='/auth/account', tags=['Cabinet Account Linking']) + + +@router.get('/linked-providers', response_model=LinkedProvidersResponse) +async def get_linked_providers( + user: User = Depends(get_current_cabinet_user), +) -> LinkedProvidersResponse: + """Return all auth methods with their link status for the current user.""" + providers: list[LinkedProvider] = [] + for provider in _ALL_PROVIDERS: + identifier = _get_provider_identifier(user, provider) + providers.append( + LinkedProvider( + provider=provider, + linked=identifier is not None, + identifier=identifier, + ) + ) + return LinkedProvidersResponse(providers=providers) + + +@router.get('/link/{provider}/init', response_model=LinkInitResponse) +async def link_provider_init( + provider: str, + user: User = Depends(get_current_cabinet_user), +) -> LinkInitResponse: + """Start OAuth flow for linking a new provider to the current account.""" + if provider not in _OAUTH_PROVIDER_COLUMNS: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail='Only OAuth providers can be linked via this endpoint', + ) + + # Check if already linked + column = _OAUTH_PROVIDER_COLUMNS[provider] + if getattr(user, column, None): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail='Provider is already linked to your account', + ) + + oauth_provider = get_provider(provider) + if not oauth_provider: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail='Requested OAuth provider is not available', + ) + + # Generate PKCE data for VK (and potentially future providers) + auth_extra = oauth_provider.prepare_auth_state() + extra_data: dict[str, str] = { + 'linking': 'true', + 'user_id': str(user.id), + } + if auth_extra: + extra_data.update(auth_extra) + + state = await generate_oauth_state(provider, extra_data=extra_data) + authorize_url = oauth_provider.get_authorization_url(state, **auth_extra) + + return LinkInitResponse(authorize_url=authorize_url, state=state) + + +@router.post('/link/{provider}/callback', response_model=LinkCallbackResponse) +async def link_provider_callback( + provider: str, + request: LinkCallbackRequest, + user: User = Depends(get_current_cabinet_user), + db: AsyncSession = Depends(get_cabinet_db), +) -> LinkCallbackResponse: + """Handle OAuth callback for linking a provider to the current account.""" + if provider not in _OAUTH_PROVIDER_COLUMNS: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail='Only OAuth providers can be linked via this endpoint', + ) + + # 1. Validate CSRF state + state_data = await validate_oauth_state(request.state, provider) + if not state_data: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail='Invalid or expired OAuth state', + ) + + # 2. Get provider instance + oauth_provider = get_provider(provider) + if not oauth_provider: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail='Requested OAuth provider is not available', + ) + + # 3. Exchange code for tokens + exchange_kwargs: dict[str, str] = {'state': request.state} + code_verifier = state_data.get('code_verifier') + if code_verifier: + exchange_kwargs['code_verifier'] = code_verifier + if request.device_id: + exchange_kwargs['device_id'] = request.device_id + + try: + token_data = await oauth_provider.exchange_code(request.code, **exchange_kwargs) + except Exception as exc: + logger.error('OAuth code exchange failed during linking', provider=provider, exc_info=exc) + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail='Failed to exchange authorization code', + ) from exc + + # 4. Fetch user info from provider + try: + user_info = await oauth_provider.get_user_info(token_data) + except Exception as exc: + logger.error('OAuth user info fetch failed during linking', provider=provider, exc_info=exc) + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail='Failed to fetch user information from provider', + ) from exc + + # 5. Check if provider_id is already linked to THIS user + column = _OAUTH_PROVIDER_COLUMNS[provider] + current_value = getattr(user, column, None) + if current_value and str(current_value) == user_info.provider_id: + return LinkCallbackResponse(success=True, message='already_linked') + + # 6. Check if provider_id is linked to ANOTHER user + existing_user = await get_user_by_oauth_provider(db, provider, user_info.provider_id) + if existing_user and existing_user.id != user.id: + # Account conflict -> create merge token + logger.info( + 'Account linking conflict: provider already linked to another user', + provider=provider, + provider_id=user_info.provider_id, + current_user_id=user.id, + existing_user_id=existing_user.id, + ) + merge_token = await create_merge_token( + primary_user_id=user.id, + secondary_user_id=existing_user.id, + provider=provider, + provider_id=user_info.provider_id, + ) + return LinkCallbackResponse( + success=False, + merge_required=True, + merge_token=merge_token, + ) + + # 7. Link the provider to current user + await set_user_oauth_provider_id(db, user, provider, user_info.provider_id) + await db.commit() + + logger.info( + 'OAuth provider linked to account', + provider=provider, + provider_id=user_info.provider_id, + user_id=user.id, + ) + return LinkCallbackResponse(success=True, message='linked') + + +@router.post('/unlink/{provider}', response_model=UnlinkResponse) +async def unlink_provider( + provider: str, + user: User = Depends(get_current_cabinet_user), + db: AsyncSession = Depends(get_cabinet_db), +) -> UnlinkResponse: + """Unlink an OAuth provider from the current account.""" + if provider not in _OAUTH_PROVIDER_COLUMNS: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail='Only OAuth providers can be unlinked via this endpoint', + ) + + column = _OAUTH_PROVIDER_COLUMNS[provider] + if not getattr(user, column, None): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail='Provider is not linked to your account', + ) + + # Ensure at least one auth method remains + if _count_auth_methods(user) <= 1: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail='Cannot unlink last authentication method', + ) + + setattr(user, column, None) + user.updated_at = datetime.now(UTC) + await db.commit() + + logger.info('OAuth provider unlinked from account', provider=provider, user_id=user.id) + return UnlinkResponse(success=True) + + +# --------------------------------------------------------------------------- +# Router 2: Merge (NO JWT required) +# --------------------------------------------------------------------------- + +merge_router = APIRouter(prefix='/auth/merge', tags=['Cabinet Account Merge']) + + +@merge_router.get('/{merge_token}', response_model=MergePreviewResponse) +async def get_merge_preview_endpoint( + merge_token: str, + db: AsyncSession = Depends(get_cabinet_db), +) -> MergePreviewResponse: + """Preview the result of merging two accounts before confirming.""" + token_data = await get_merge_token_data(merge_token) + if not token_data: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail='Merge token is invalid or expired', + ) + + primary_user_id: int = token_data['primary_user_id'] + secondary_user_id: int = token_data['secondary_user_id'] + + try: + preview = await get_merge_preview(db, primary_user_id, secondary_user_id) + except ValueError as exc: + logger.error('Merge preview failed', error=str(exc)) + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail='One or both users not found', + ) from exc + + # Calculate remaining TTL + created_at_str: str = token_data.get('created_at', '') + try: + created_at = datetime.fromisoformat(created_at_str) + if created_at.tzinfo is None: + created_at = created_at.replace(tzinfo=UTC) + elapsed = (datetime.now(UTC) - created_at).total_seconds() + expires_in_seconds = max(0, int(MERGE_TOKEN_TTL_SECONDS - elapsed)) + except (ValueError, TypeError): + expires_in_seconds = 0 + + return MergePreviewResponse( + primary=MergePreviewUser(**preview['primary']), + secondary=MergePreviewUser(**preview['secondary']), + expires_in_seconds=expires_in_seconds, + ) + + +@merge_router.post('/{merge_token}', response_model=MergeResponse) +async def execute_merge_endpoint( + merge_token: str, + request: MergeRequest, + db: AsyncSession = Depends(get_cabinet_db), +) -> MergeResponse: + """Execute account merge. Consumes the merge token (one-time use).""" + # 1. Consume token atomically + token_data = await consume_merge_token(merge_token) + if not token_data: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail='Merge token is invalid, expired, or already consumed', + ) + + primary_user_id: int = token_data['primary_user_id'] + secondary_user_id: int = token_data['secondary_user_id'] + provider: str = token_data.get('provider', '') + provider_id: str = token_data.get('provider_id', '') + + # 2. Validate keep_subscription_from is one of the two user IDs + if request.keep_subscription_from not in (primary_user_id, secondary_user_id): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail='keep_subscription_from must be one of the two user IDs being merged', + ) + + # Convert user_id to 'primary'/'secondary' string for execute_merge() + keep_from: str = 'primary' if request.keep_subscription_from == primary_user_id else 'secondary' + + # 3. Execute merge + try: + merged_user = await execute_merge( + db=db, + primary_user_id=primary_user_id, + secondary_user_id=secondary_user_id, + keep_subscription_from=keep_from, + provider=provider, + provider_id=provider_id, + ) + await db.commit() + except ValueError as exc: + await db.rollback() + logger.error('Merge execution failed (ValueError)', error=str(exc)) + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=str(exc), + ) from exc + except Exception as exc: + await db.rollback() + logger.error('Merge execution failed', exc_info=exc) + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail='Account merge failed due to an internal error', + ) from exc + + # 4. Re-fetch merged user with full relationships for auth response + merged_user = await get_user_by_id(db, primary_user_id) + if not merged_user: + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail='Failed to load merged user', + ) + + # 5. Create auth tokens for the merged user + try: + auth_response = await _create_auth_response(merged_user, db) + await _store_refresh_token(db, merged_user.id, auth_response.refresh_token, device_info='merge') + except Exception as exc: + logger.error('Failed to create auth tokens after merge', exc_info=exc) + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail='Merge succeeded but failed to create new session', + ) from exc + + logger.info( + 'Account merge completed successfully', + primary_user_id=primary_user_id, + secondary_user_id=secondary_user_id, + provider=provider, + ) + + return MergeResponse( + success=True, + access_token=auth_response.access_token, + refresh_token=auth_response.refresh_token, + user=_user_to_response(merged_user), + ) diff --git a/app/services/account_merge_service.py b/app/services/account_merge_service.py new file mode 100644 index 00000000..98840608 --- /dev/null +++ b/app/services/account_merge_service.py @@ -0,0 +1,483 @@ +from __future__ import annotations + +from contextlib import asynccontextmanager +from datetime import UTC, datetime +from typing import Any + +import structlog +from sqlalchemy import update +from sqlalchemy.ext.asyncio import AsyncSession + +from app.config import settings +from app.database.crud.user import get_user_by_id +from app.database.models import ( + CabinetRefreshToken, + CloudPaymentsPayment, + CryptoBotPayment, + FreekassaPayment, + HeleketPayment, + KassaAiPayment, + MulenPayPayment, + Pal24Payment, + PartnerStatus, + PlategaPayment, + ReferralEarning, + Subscription, + Transaction, + User, + UserStatus, + WataPayment, + WithdrawalRequest, + YooKassaPayment, +) +from app.external.remnawave_api import RemnaWaveAPI + + +logger = structlog.get_logger(__name__) + +# OAuth-поля, которые можно перенести между аккаунтами +_OAUTH_FIELDS: tuple[str, ...] = ('google_id', 'yandex_id', 'discord_id', 'vk_id') + +# Все платёжные таблицы с колонкой user_id +_PAYMENT_MODELS: tuple[type, ...] = ( + YooKassaPayment, + CryptoBotPayment, + HeleketPayment, + MulenPayPayment, + Pal24Payment, + WataPayment, + PlategaPayment, + CloudPaymentsPayment, + FreekassaPayment, + KassaAiPayment, +) + +# Приоритет партнёрских статусов (чем выше число — тем приоритетнее) +_PARTNER_STATUS_PRIORITY: dict[str, int] = { + PartnerStatus.NONE.value: 0, + PartnerStatus.PENDING.value: 1, + PartnerStatus.REJECTED.value: 2, + PartnerStatus.APPROVED.value: 3, +} + + +def _compute_auth_methods(user: User) -> list[str]: + """Вычисляет список методов авторизации пользователя.""" + methods: list[str] = [] + if user.telegram_id: + methods.append('telegram') + if user.email and user.password_hash: + methods.append('email') + if user.google_id: + methods.append('google') + if user.yandex_id: + methods.append('yandex') + if user.discord_id: + methods.append('discord') + if user.vk_id: + methods.append('vk') + return methods + + +def _build_subscription_preview(sub: Subscription | None) -> dict[str, Any] | None: + """Формирует превью данных подписки.""" + if sub is None: + return None + tariff_name: str | None = None + if sub.tariff: + tariff_name = sub.tariff.name + return { + 'status': sub.status, + 'is_trial': sub.is_trial, + 'end_date': sub.end_date, + 'traffic_limit_gb': sub.traffic_limit_gb, + 'traffic_used_gb': sub.traffic_used_gb, + 'device_limit': sub.device_limit, + 'tariff_name': tariff_name, + 'autopay_enabled': sub.autopay_enabled, + } + + +def _build_user_preview(user: User) -> dict[str, Any]: + """Формирует превью данных пользователя для предварительного просмотра мержа.""" + return { + 'id': user.id, + 'username': user.username, + 'first_name': user.first_name, + 'email': user.email, + 'auth_methods': _compute_auth_methods(user), + 'balance_kopeks': user.balance_kopeks, + 'subscription': _build_subscription_preview(user.subscription), + 'created_at': user.created_at, + } + + +async def get_merge_preview( + db: AsyncSession, + primary_user_id: int, + secondary_user_id: int, +) -> dict[str, Any]: + """Возвращает превью данных обоих аккаунтов для подтверждения мержа. + + Args: + db: Сессия БД. + primary_user_id: ID основного аккаунта (останется). + secondary_user_id: ID вторичного аккаунта (будет поглощён). + + Returns: + Словарь с ключами 'primary' и 'secondary', содержащими превью данных. + + Raises: + ValueError: Если один из пользователей не найден или совпадают. + """ + if primary_user_id == secondary_user_id: + raise ValueError('primary_user_id и secondary_user_id не могут совпадать') + + primary = await get_user_by_id(db, primary_user_id) + secondary = await get_user_by_id(db, secondary_user_id) + + if not primary: + raise ValueError(f'Основной пользователь (id={primary_user_id}) не найден') + if not secondary: + raise ValueError(f'Вторичный пользователь (id={secondary_user_id}) не найден') + + return { + 'primary': _build_user_preview(primary), + 'secondary': _build_user_preview(secondary), + } + + +@asynccontextmanager +async def _get_remnawave_api() -> RemnaWaveAPI: + """Создаёт экземпляр RemnaWave API клиента (паттерн из RemnaWaveService).""" + auth_params = settings.get_remnawave_auth_params() + base_url = (auth_params.get('base_url') or '').strip() + api_key = (auth_params.get('api_key') or '').strip() + + if not base_url or not api_key: + raise RuntimeError('RemnaWave API не настроен (REMNAWAVE_API_URL / REMNAWAVE_API_KEY)') + + api = RemnaWaveAPI( + base_url=base_url, + api_key=api_key, + secret_key=auth_params.get('secret_key'), + username=auth_params.get('username'), + password=auth_params.get('password'), + caddy_token=auth_params.get('caddy_token'), + auth_type=auth_params.get('auth_type') or 'api_key', + ) + async with api: + yield api + + +async def _delete_remnawave_user_with_fallback(remnawave_uuid: str) -> None: + """Удаляет пользователя из RemnaWave. При неудаче — деактивирует как fallback.""" + try: + async with _get_remnawave_api() as api: + deleted = await api.delete_user(remnawave_uuid) + if deleted: + logger.info( + 'RemnaWave пользователь удалён при мерже', + remnawave_uuid=remnawave_uuid, + ) + else: + logger.warning( + 'RemnaWave delete_user вернул False, пробуем disable', + remnawave_uuid=remnawave_uuid, + ) + await api.disable_user(remnawave_uuid) + logger.info( + 'RemnaWave пользователь деактивирован как fallback при мерже', + remnawave_uuid=remnawave_uuid, + ) + except Exception as exc: + logger.warning( + 'Не удалось удалить RemnaWave пользователя, пробуем disable', + remnawave_uuid=remnawave_uuid, + error=exc, + ) + try: + async with _get_remnawave_api() as api: + await api.disable_user(remnawave_uuid) + logger.info( + 'RemnaWave пользователь деактивирован как fallback при мерже', + remnawave_uuid=remnawave_uuid, + ) + except Exception as fallback_exc: + logger.error( + 'Не удалось ни удалить, ни деактивировать RemnaWave пользователя', + remnawave_uuid=remnawave_uuid, + error=fallback_exc, + ) + + +async def _handle_subscription_merge( + db: AsyncSession, + primary: User, + secondary: User, + keep_subscription_from: str, +) -> None: + """Обрабатывает мерж подписок между двумя аккаунтами. + + Args: + db: Сессия БД. + primary: Основной пользователь. + secondary: Вторичный пользователь. + keep_subscription_from: 'primary' или 'secondary' — чью подписку оставить. + """ + primary_sub = primary.subscription + secondary_sub = secondary.subscription + has_primary_sub = primary_sub is not None + has_secondary_sub = secondary_sub is not None + + # Ни у кого нет подписки — ничего не делаем + if not has_primary_sub and not has_secondary_sub: + logger.info( + 'Мерж подписок: ни у кого нет подписки', + primary_id=primary.id, + secondary_id=secondary.id, + ) + return + + # Подписка только у primary — удаляем RemnaWave юзера secondary (если есть) + if has_primary_sub and not has_secondary_sub: + if secondary.remnawave_uuid: + await _delete_remnawave_user_with_fallback(secondary.remnawave_uuid) + secondary.remnawave_uuid = None + logger.info( + 'Мерж подписок: оставлена подписка primary, secondary не имел подписки', + primary_id=primary.id, + secondary_id=secondary.id, + ) + return + + # Подписка только у secondary — переносим на primary + if not has_primary_sub and has_secondary_sub: + assert secondary_sub is not None + secondary_sub.user_id = primary.id + # Переносим remnawave_uuid с secondary на primary + if secondary.remnawave_uuid: + primary.remnawave_uuid = secondary.remnawave_uuid + secondary.remnawave_uuid = None + logger.info( + 'Мерж подписок: перенесена подписка secondary на primary', + primary_id=primary.id, + secondary_id=secondary.id, + ) + return + + # Обе подписки есть — выбираем по keep_subscription_from + assert primary_sub is not None + assert secondary_sub is not None + + if keep_subscription_from == 'secondary': + # Удаляем подписку primary из RemnaWave + if primary.remnawave_uuid: + await _delete_remnawave_user_with_fallback(primary.remnawave_uuid) + primary.remnawave_uuid = None + # Удаляем запись подписки primary + await db.delete(primary_sub) + await db.flush() + # Переносим подписку secondary на primary + secondary_sub.user_id = primary.id + # Переносим remnawave_uuid + if secondary.remnawave_uuid: + primary.remnawave_uuid = secondary.remnawave_uuid + secondary.remnawave_uuid = None + logger.info( + 'Мерж подписок: оставлена подписка secondary, подписка primary удалена', + primary_id=primary.id, + secondary_id=secondary.id, + ) + else: + # keep_subscription_from == 'primary' (по умолчанию) + # Удаляем подписку secondary из RemnaWave + if secondary.remnawave_uuid: + await _delete_remnawave_user_with_fallback(secondary.remnawave_uuid) + secondary.remnawave_uuid = None + # Удаляем запись подписки secondary + await db.delete(secondary_sub) + logger.info( + 'Мерж подписок: оставлена подписка primary, подписка secondary удалена', + primary_id=primary.id, + secondary_id=secondary.id, + ) + + +async def execute_merge( + db: AsyncSession, + primary_user_id: int, + secondary_user_id: int, + keep_subscription_from: str = 'primary', + provider: str | None = None, + provider_id: str | None = None, +) -> User: + """Выполняет атомарный мерж двух аккаунтов. Caller отвечает за commit/rollback. + + Переносит все данные с secondary на primary, помечает secondary как deleted. + + Args: + db: Сессия БД (caller управляет транзакцией). + primary_user_id: ID основного аккаунта. + secondary_user_id: ID вторичного аккаунта. + keep_subscription_from: 'primary' или 'secondary' — чью подписку оставить. + provider: OAuth-провайдер, инициировавший мерж (для логирования). + provider_id: ID провайдера (для логирования). + + Returns: + Обновлённый объект primary User. + + Raises: + ValueError: Если пользователь не найден, совпадают ID, или secondary уже удалён. + """ + if primary_user_id == secondary_user_id: + raise ValueError('primary_user_id и secondary_user_id не могут совпадать') + + primary = await get_user_by_id(db, primary_user_id) + secondary = await get_user_by_id(db, secondary_user_id) + + if not primary: + raise ValueError(f'Основной пользователь (id={primary_user_id}) не найден') + if not secondary: + raise ValueError(f'Вторичный пользователь (id={secondary_user_id}) не найден') + if secondary.status == UserStatus.DELETED.value: + raise ValueError(f'Вторичный пользователь (id={secondary_user_id}) уже удалён') + + logger.info( + 'Начинаем мерж аккаунтов', + primary_id=primary.id, + secondary_id=secondary.id, + keep_subscription_from=keep_subscription_from, + provider=provider, + provider_id=provider_id, + ) + + # 1. Перенос OAuth ID + for field in _OAUTH_FIELDS: + secondary_value = getattr(secondary, field) + primary_value = getattr(primary, field) + if secondary_value and not primary_value: + setattr(primary, field, secondary_value) + setattr(secondary, field, None) + logger.info( + 'Перенесён OAuth ID', + field=field, + primary_id=primary.id, + secondary_id=secondary.id, + ) + + # 2. Перенос telegram_id + if secondary.telegram_id and not primary.telegram_id: + primary.telegram_id = secondary.telegram_id + secondary.telegram_id = None + logger.info( + 'Перенесён telegram_id', + primary_id=primary.id, + secondary_id=secondary.id, + ) + + # 3. Перенос email + password + if not primary.email and secondary.email: + primary.email = secondary.email + primary.email_verified = secondary.email_verified + primary.email_verified_at = secondary.email_verified_at + primary.password_hash = secondary.password_hash + # Очищаем на secondary для освобождения unique constraint + secondary.email = None + secondary.email_verified = False + secondary.email_verified_at = None + secondary.password_hash = None + logger.info( + 'Перенесены email и пароль', + primary_id=primary.id, + secondary_id=secondary.id, + ) + + # 4. Суммируем баланс + if secondary.balance_kopeks > 0: + primary.balance_kopeks += secondary.balance_kopeks + logger.info( + 'Перенесён баланс', + primary_id=primary.id, + secondary_id=secondary.id, + transferred_kopeks=secondary.balance_kopeks, + ) + secondary.balance_kopeks = 0 + + # 5. Мерж подписок + await _handle_subscription_merge(db, primary, secondary, keep_subscription_from) + + # 6. Переназначение транзакций + await db.execute(update(Transaction).where(Transaction.user_id == secondary.id).values(user_id=primary.id)) + + # 7. Переназначение всех платёжных таблиц + for payment_model in _PAYMENT_MODELS: + await db.execute(update(payment_model).where(payment_model.user_id == secondary.id).values(user_id=primary.id)) + + # 8. Переназначение referral_earnings (обе колонки) + await db.execute(update(ReferralEarning).where(ReferralEarning.user_id == secondary.id).values(user_id=primary.id)) + await db.execute( + update(ReferralEarning).where(ReferralEarning.referral_id == secondary.id).values(referral_id=primary.id) + ) + + # 9. Переназначение реферальной цепочки + await db.execute(update(User).where(User.referred_by_id == secondary.id).values(referred_by_id=primary.id)) + + # 10. Переназначение withdrawal_requests + await db.execute( + update(WithdrawalRequest).where(WithdrawalRequest.user_id == secondary.id).values(user_id=primary.id) + ) + + # 11. Инвалидация refresh-токенов secondary + now = datetime.now(UTC) + await db.execute( + update(CabinetRefreshToken) + .where( + CabinetRefreshToken.user_id == secondary.id, + CabinetRefreshToken.revoked_at.is_(None), + ) + .values(revoked_at=now) + ) + + # 12. Перенос partner_status (оставляем более приоритетный) + primary_priority = _PARTNER_STATUS_PRIORITY.get(primary.partner_status, 0) + secondary_priority = _PARTNER_STATUS_PRIORITY.get(secondary.partner_status, 0) + if secondary_priority > primary_priority: + primary.partner_status = secondary.partner_status + logger.info( + 'Перенесён partner_status', + primary_id=primary.id, + secondary_id=secondary.id, + new_status=primary.partner_status, + ) + + # 13. Перенос referral_commission_percent + if secondary.referral_commission_percent is not None and primary.referral_commission_percent is None: + primary.referral_commission_percent = secondary.referral_commission_percent + logger.info( + 'Перенесён referral_commission_percent', + primary_id=primary.id, + secondary_id=secondary.id, + value=primary.referral_commission_percent, + ) + + # 14. Помечаем secondary как удалённый + secondary.status = UserStatus.DELETED.value + secondary.referral_code = None + secondary.remnawave_uuid = None + # email уже очищен выше если был перенесён, иначе очищаем для unique constraint + if secondary.email: + secondary.email = None + secondary.updated_at = now + + logger.info( + 'Мерж аккаунтов завершён', + primary_id=primary.id, + secondary_id=secondary.id, + provider=provider, + ) + + # 15. flush (не commit — caller управляет транзакцией) + await db.flush() + + return primary diff --git a/tests/services/test_account_merge_service.py b/tests/services/test_account_merge_service.py new file mode 100644 index 00000000..ccfcaf8b --- /dev/null +++ b/tests/services/test_account_merge_service.py @@ -0,0 +1,643 @@ +"""Tests for app.services.account_merge_service.""" + +from datetime import UTC, datetime +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +import pytest + +from app.services import account_merge_service +from app.services.account_merge_service import ( + _build_subscription_preview, + _build_user_preview, + _compute_auth_methods, + execute_merge, + get_merge_preview, +) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_user( + *, + id: int = 1, + telegram_id: int | None = None, + email: str | None = None, + email_verified: bool = False, + email_verified_at: datetime | None = None, + password_hash: str | None = None, + google_id: str | None = None, + yandex_id: str | None = None, + discord_id: str | None = None, + vk_id: int | None = None, + balance_kopeks: int = 0, + username: str | None = None, + first_name: str | None = None, + status: str = 'active', + partner_status: str = 'none', + referral_code: str | None = None, + referral_commission_percent: int | None = None, + remnawave_uuid: str | None = None, + subscription: object | None = None, + created_at: datetime | None = None, + updated_at: datetime | None = None, +) -> SimpleNamespace: + return SimpleNamespace( + id=id, + telegram_id=telegram_id, + email=email, + email_verified=email_verified, + email_verified_at=email_verified_at, + password_hash=password_hash, + google_id=google_id, + yandex_id=yandex_id, + discord_id=discord_id, + vk_id=vk_id, + balance_kopeks=balance_kopeks, + username=username, + first_name=first_name, + status=status, + partner_status=partner_status, + referral_code=referral_code, + referral_commission_percent=referral_commission_percent, + remnawave_uuid=remnawave_uuid, + subscription=subscription, + created_at=created_at or datetime(2024, 1, 1, tzinfo=UTC), + updated_at=updated_at or datetime(2024, 1, 1, tzinfo=UTC), + ) + + +def _make_subscription( + *, + user_id: int = 1, + status: str = 'active', + is_trial: bool = False, + end_date: datetime | None = None, + traffic_limit_gb: float = 100.0, + traffic_used_gb: float = 10.0, + device_limit: int = 3, + tariff_name: str = 'Basic', + autopay_enabled: bool = False, +) -> SimpleNamespace: + tariff = SimpleNamespace(name=tariff_name) + return SimpleNamespace( + user_id=user_id, + status=status, + is_trial=is_trial, + end_date=end_date or datetime(2025, 1, 1, tzinfo=UTC), + traffic_limit_gb=traffic_limit_gb, + traffic_used_gb=traffic_used_gb, + device_limit=device_limit, + tariff=tariff, + autopay_enabled=autopay_enabled, + ) + + +def _make_db() -> SimpleNamespace: + return SimpleNamespace( + execute=AsyncMock(), + delete=AsyncMock(), + flush=AsyncMock(), + ) + + +# --------------------------------------------------------------------------- +# _compute_auth_methods +# --------------------------------------------------------------------------- + + +class TestComputeAuthMethods: + def test_no_methods(self): + user = _make_user() + assert _compute_auth_methods(user) == [] + + def test_telegram_only(self): + user = _make_user(telegram_id=12345) + assert _compute_auth_methods(user) == ['telegram'] + + def test_email_only(self): + user = _make_user(email='test@example.com', password_hash='hash123') + assert _compute_auth_methods(user) == ['email'] + + def test_email_without_password_not_counted(self): + user = _make_user(email='test@example.com') + assert _compute_auth_methods(user) == [] + + def test_all_methods(self): + user = _make_user( + telegram_id=12345, + email='test@example.com', + password_hash='hash', + google_id='g123', + yandex_id='y123', + discord_id='d123', + vk_id=99999, + ) + assert _compute_auth_methods(user) == ['telegram', 'email', 'google', 'yandex', 'discord', 'vk'] + + def test_oauth_only(self): + user = _make_user(google_id='g123', discord_id='d123') + assert _compute_auth_methods(user) == ['google', 'discord'] + + +# --------------------------------------------------------------------------- +# _build_subscription_preview +# --------------------------------------------------------------------------- + + +class TestBuildSubscriptionPreview: + def test_none_subscription(self): + assert _build_subscription_preview(None) is None + + def test_valid_subscription(self): + sub = _make_subscription(tariff_name='Premium') + result = _build_subscription_preview(sub) + assert result['tariff_name'] == 'Premium' + assert result['status'] == 'active' + assert result['is_trial'] is False + assert result['device_limit'] == 3 + + def test_subscription_without_tariff(self): + sub = _make_subscription() + sub.tariff = None + result = _build_subscription_preview(sub) + assert result['tariff_name'] is None + + +# --------------------------------------------------------------------------- +# _build_user_preview +# --------------------------------------------------------------------------- + + +class TestBuildUserPreview: + def test_basic_user(self): + user = _make_user(id=42, username='alice', email='a@b.com', balance_kopeks=5000) + result = _build_user_preview(user) + assert result['id'] == 42 + assert result['username'] == 'alice' + assert result['balance_kopeks'] == 5000 + assert result['subscription'] is None + + def test_user_with_subscription(self): + sub = _make_subscription(user_id=1) + user = _make_user(id=1, subscription=sub) + result = _build_user_preview(user) + assert result['subscription'] is not None + assert result['subscription']['status'] == 'active' + + +# --------------------------------------------------------------------------- +# get_merge_preview +# --------------------------------------------------------------------------- + + +class TestGetMergePreview: + async def test_same_user_ids_raises(self): + db = _make_db() + with pytest.raises(ValueError, match='не могут совпадать'): + await get_merge_preview(db, 1, 1) + + async def test_primary_not_found_raises(self, monkeypatch): + db = _make_db() + monkeypatch.setattr(account_merge_service, 'get_user_by_id', AsyncMock(return_value=None)) + with pytest.raises(ValueError, match='Основной пользователь'): + await get_merge_preview(db, 1, 2) + + async def test_secondary_not_found_raises(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1) + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, None]), + ) + with pytest.raises(ValueError, match='Вторичный пользователь'): + await get_merge_preview(db, 1, 2) + + async def test_success(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1, username='primary', telegram_id=111) + secondary = _make_user(id=2, username='secondary', google_id='g123') + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + result = await get_merge_preview(db, 1, 2) + assert result['primary']['id'] == 1 + assert result['secondary']['id'] == 2 + assert 'telegram' in result['primary']['auth_methods'] + assert 'google' in result['secondary']['auth_methods'] + + +# --------------------------------------------------------------------------- +# execute_merge — validation +# --------------------------------------------------------------------------- + + +class TestExecuteMergeValidation: + async def test_same_ids_raises(self): + db = _make_db() + with pytest.raises(ValueError, match='не могут совпадать'): + await execute_merge(db, 1, 1) + + async def test_primary_not_found_raises(self, monkeypatch): + db = _make_db() + monkeypatch.setattr(account_merge_service, 'get_user_by_id', AsyncMock(return_value=None)) + with pytest.raises(ValueError, match='Основной пользователь'): + await execute_merge(db, 1, 2) + + async def test_secondary_not_found_raises(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1) + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, None]), + ) + with pytest.raises(ValueError, match='Вторичный пользователь'): + await execute_merge(db, 1, 2) + + async def test_deleted_secondary_raises(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1) + secondary = _make_user(id=2, status='deleted') + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with pytest.raises(ValueError, match='уже удалён'): + await execute_merge(db, 1, 2) + + +# --------------------------------------------------------------------------- +# execute_merge — data transfer +# --------------------------------------------------------------------------- + + +def _patch_remnawave_delete(): + return patch.object( + account_merge_service, + '_delete_remnawave_user_with_fallback', + new_callable=AsyncMock, + ) + + +class TestExecuteMergeOAuthTransfer: + async def test_transfers_oauth_ids(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1, google_id='g_primary') + secondary = _make_user(id=2, yandex_id='y_sec', discord_id='d_sec', vk_id=12345) + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + result = await execute_merge(db, 1, 2) + + # google_id stays on primary (already set) + assert result.google_id == 'g_primary' + # transferred from secondary + assert result.yandex_id == 'y_sec' + assert result.discord_id == 'd_sec' + assert result.vk_id == 12345 + # cleared on secondary + assert secondary.yandex_id is None + assert secondary.discord_id is None + assert secondary.vk_id is None + + async def test_does_not_overwrite_existing_oauth(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1, google_id='g_primary') + secondary = _make_user(id=2, google_id='g_secondary') + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + result = await execute_merge(db, 1, 2) + + # Primary keeps its own google_id + assert result.google_id == 'g_primary' + + +class TestExecuteMergeTelegramTransfer: + async def test_transfers_telegram_id(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1) + secondary = _make_user(id=2, telegram_id=99999) + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + result = await execute_merge(db, 1, 2) + + assert result.telegram_id == 99999 + assert secondary.telegram_id is None + + async def test_does_not_overwrite_telegram_id(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1, telegram_id=11111) + secondary = _make_user(id=2, telegram_id=22222) + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + result = await execute_merge(db, 1, 2) + + assert result.telegram_id == 11111 + + +class TestExecuteMergeEmailTransfer: + async def test_transfers_email_and_password(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1) + secondary = _make_user( + id=2, + email='sec@example.com', + email_verified=True, + email_verified_at=datetime(2024, 6, 1, tzinfo=UTC), + password_hash='hash_sec', + ) + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + result = await execute_merge(db, 1, 2) + + assert result.email == 'sec@example.com' + assert result.email_verified is True + assert result.password_hash == 'hash_sec' + # secondary cleared + assert secondary.email is None + assert secondary.password_hash is None + + async def test_does_not_overwrite_existing_email(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1, email='pri@example.com', password_hash='hash_pri') + secondary = _make_user(id=2, email='sec@example.com', password_hash='hash_sec') + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + result = await execute_merge(db, 1, 2) + + assert result.email == 'pri@example.com' + + +class TestExecuteMergeBalance: + async def test_sums_balances(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1, balance_kopeks=5000) + secondary = _make_user(id=2, balance_kopeks=3000) + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + result = await execute_merge(db, 1, 2) + + assert result.balance_kopeks == 8000 + assert secondary.balance_kopeks == 0 + + async def test_zero_secondary_balance_unchanged(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1, balance_kopeks=5000) + secondary = _make_user(id=2, balance_kopeks=0) + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + result = await execute_merge(db, 1, 2) + + assert result.balance_kopeks == 5000 + + +class TestExecuteMergePartnerStatus: + async def test_higher_priority_transferred(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1, partner_status='none') + secondary = _make_user(id=2, partner_status='approved') + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + result = await execute_merge(db, 1, 2) + + assert result.partner_status == 'approved' + + async def test_lower_priority_not_overwritten(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1, partner_status='approved') + secondary = _make_user(id=2, partner_status='pending') + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + result = await execute_merge(db, 1, 2) + + assert result.partner_status == 'approved' + + +class TestExecuteMergeReferralCommission: + async def test_transfers_if_primary_has_none(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1, referral_commission_percent=None) + secondary = _make_user(id=2, referral_commission_percent=15) + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + result = await execute_merge(db, 1, 2) + + assert result.referral_commission_percent == 15 + + async def test_does_not_overwrite_if_primary_has_value(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1, referral_commission_percent=20) + secondary = _make_user(id=2, referral_commission_percent=15) + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + result = await execute_merge(db, 1, 2) + + assert result.referral_commission_percent == 20 + + +# --------------------------------------------------------------------------- +# execute_merge — secondary marked as deleted +# --------------------------------------------------------------------------- + + +class TestExecuteMergeSecondaryDeleted: + async def test_secondary_marked_deleted(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1) + secondary = _make_user(id=2, referral_code='REF123', email='sec@e.com') + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + await execute_merge(db, 1, 2) + + assert secondary.status == 'deleted' + assert secondary.referral_code is None + assert secondary.remnawave_uuid is None + assert secondary.email is None + + async def test_db_flush_called(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1) + secondary = _make_user(id=2) + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + await execute_merge(db, 1, 2) + + db.flush.assert_awaited_once() + + +# --------------------------------------------------------------------------- +# execute_merge — subscription merge scenarios +# --------------------------------------------------------------------------- + + +class TestExecuteMergeSubscription: + async def test_neither_has_subscription(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1) + secondary = _make_user(id=2) + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete() as mock_del: + await execute_merge(db, 1, 2) + mock_del.assert_not_awaited() + + async def test_only_primary_has_subscription(self, monkeypatch): + db = _make_db() + sub = _make_subscription(user_id=1) + primary = _make_user(id=1, subscription=sub, remnawave_uuid='rw-primary') + secondary = _make_user(id=2, remnawave_uuid='rw-secondary') + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete() as mock_del: + await execute_merge(db, 1, 2) + mock_del.assert_awaited_once_with('rw-secondary') + + # secondary remnawave_uuid cleared + assert secondary.remnawave_uuid is None + + async def test_only_secondary_has_subscription(self, monkeypatch): + db = _make_db() + sub = _make_subscription(user_id=2) + primary = _make_user(id=1) + secondary = _make_user(id=2, subscription=sub, remnawave_uuid='rw-secondary') + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + await execute_merge(db, 1, 2) + + # Subscription transferred to primary + assert sub.user_id == 1 + assert primary.remnawave_uuid == 'rw-secondary' + assert secondary.remnawave_uuid is None + + async def test_both_have_subscription_keep_primary(self, monkeypatch): + db = _make_db() + sub_p = _make_subscription(user_id=1) + sub_s = _make_subscription(user_id=2) + primary = _make_user(id=1, subscription=sub_p, remnawave_uuid='rw-primary') + secondary = _make_user(id=2, subscription=sub_s, remnawave_uuid='rw-secondary') + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete() as mock_del: + await execute_merge(db, 1, 2, keep_subscription_from='primary') + mock_del.assert_awaited_once_with('rw-secondary') + + db.delete.assert_awaited_once_with(sub_s) + + async def test_both_have_subscription_keep_secondary(self, monkeypatch): + db = _make_db() + sub_p = _make_subscription(user_id=1) + sub_s = _make_subscription(user_id=2) + primary = _make_user(id=1, subscription=sub_p, remnawave_uuid='rw-primary') + secondary = _make_user(id=2, subscription=sub_s, remnawave_uuid='rw-secondary') + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete() as mock_del: + await execute_merge(db, 1, 2, keep_subscription_from='secondary') + mock_del.assert_awaited_once_with('rw-primary') + + db.delete.assert_awaited_once_with(sub_p) + # Secondary subscription transferred + assert sub_s.user_id == 1 + assert primary.remnawave_uuid == 'rw-secondary' + + +# --------------------------------------------------------------------------- +# execute_merge — bulk updates called +# --------------------------------------------------------------------------- + + +class TestExecuteMergeBulkUpdates: + async def test_execute_called_for_transactions_and_payments(self, monkeypatch): + db = _make_db() + primary = _make_user(id=1) + secondary = _make_user(id=2) + monkeypatch.setattr( + account_merge_service, + 'get_user_by_id', + AsyncMock(side_effect=[primary, secondary]), + ) + with _patch_remnawave_delete(): + await execute_merge(db, 1, 2) + + # Transaction + 10 payment models + 2 referral_earnings + 1 referral chain + # + 1 withdrawal_requests + 1 refresh tokens = 16 total execute calls + assert db.execute.await_count == 16