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:
@@ -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
|
||||||
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
from fastapi import APIRouter
|
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_apps import router as admin_apps_router
|
||||||
from .admin_audit_log import router as admin_audit_log_router
|
from .admin_audit_log import router as admin_audit_log_router
|
||||||
from .admin_ban_system import router as admin_ban_system_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
|
# Include all sub-routers
|
||||||
router.include_router(auth_router)
|
router.include_router(auth_router)
|
||||||
router.include_router(oauth_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(subscription_router)
|
||||||
router.include_router(balance_router)
|
router.include_router(balance_router)
|
||||||
router.include_router(referral_router)
|
router.include_router(referral_router)
|
||||||
|
|||||||
@@ -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),
|
||||||
|
)
|
||||||
@@ -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
|
||||||
Reference in New Issue
Block a user