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.
This commit is contained in:
Fringg
2026-03-04 07:24:15 +03:00
parent 57aaca82f5
commit dc7b8dc72a
5 changed files with 1719 additions and 0 deletions
+101
View File
@@ -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
+3
View File
@@ -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)
+489
View File
@@ -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),
)
+483
View File
@@ -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
@@ -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