From 5e47a1496bbd4d0e42d8bd90f2a2da377971279f Mon Sep 17 00:00:00 2001 From: 3252a8 <3252a8@proton.me> Date: Tue, 2 Jun 2026 00:13:54 +0300 Subject: [PATCH] feat: add Remnashop migration import Add the Remnashop legacy importer, compatibility tables, admin toggles, referral and promo lookup compatibility, and tests for the migration flow. --- .../bot/app/web/admin_settings_manifest.py | 33 + backend/bot/app/web/webapp/auth.py | 92 +- backend/bot/handlers/user/start.py | 89 +- backend/bot/services/promo_code_service.py | 18 +- backend/config/settings.py | 18 + backend/db/dal/promo_code_dal.py | 48 +- backend/db/dal/user_dal.py | 99 +- backend/db/migrator.py | 78 ++ backend/db/models.py | 33 +- backend/scripts/__init__.py | 2 + backend/scripts/import_legacy.py | 1164 +++++++++++++++++ .../src/admin/sections/SettingsSection.svelte | 1 + locales/en.json | 10 + locales/ru.json | 10 + tests/test_admin_settings_manifest_i18n.py | 28 + tests/test_remnashop_import.py | 44 + tests/test_support_migration.py | 23 +- tests/test_user_dal.py | 4 + 18 files changed, 1731 insertions(+), 63 deletions(-) create mode 100644 backend/scripts/__init__.py create mode 100644 backend/scripts/import_legacy.py create mode 100644 tests/test_remnashop_import.py diff --git a/backend/bot/app/web/admin_settings_manifest.py b/backend/bot/app/web/admin_settings_manifest.py index 9c86789..cd8680e 100644 --- a/backend/bot/app/web/admin_settings_manifest.py +++ b/backend/bot/app/web/admin_settings_manifest.py @@ -402,6 +402,38 @@ SETTINGS_MANIFEST: List[SettingField] = [ "REFERRAL_WELCOME_BONUS_DAYS", "int", "referral", "Приветственный бонус (дней)", min=0 ), SettingField("LEGACY_REFS", "bool", "referral", "Поддержка старых ref-ссылок"), + SettingField( + "MIGRATION_REMNASHOP_REFERRAL_CODE_COMPAT_ENABLED", + "bool", + "migrations", + "Старые ref-ссылки Remnashop", + "Принимать импортированные ref-коды Remnashop вместе с текущими кодами пользователей.", + subsection="Remnashop", + ), + SettingField( + "MIGRATION_REMNASHOP_PROMO_CODE_COMPAT_ENABLED", + "bool", + "migrations", + "Старые промокоды Remnashop", + "Пробовать точное совпадение промокода перед обычной uppercase-нормализацией.", + subsection="Remnashop", + ), + SettingField( + "MIGRATION_REMNASHOP_IMPORTED_AT", + "string", + "migrations", + "Последний импорт Remnashop", + "Заполняется скриптом импорта. Можно очистить, если отметка больше не нужна.", + subsection="Remnashop", + ), + SettingField( + "MIGRATION_REMNASHOP_NOTES", + "text", + "migrations", + "Заметки по миграции Remnashop", + "Внутренние заметки оператора по перенесенному инстансу.", + subsection="Remnashop", + ), # ─── Notifications ───────────────────────────────────────────── SettingField( "SUBSCRIPTION_NOTIFICATIONS_ENABLED", @@ -734,6 +766,7 @@ def manifest_payload() -> List[dict]: "devices": 10, "subscription_guides": 10, "system": 12, + "migrations": 13, } exclusive_map = { key: opposite diff --git a/backend/bot/app/web/webapp/auth.py b/backend/bot/app/web/webapp/auth.py index f1bfbe5..e1a7f2c 100644 --- a/backend/bot/app/web/webapp/auth.py +++ b/backend/bot/app/web/webapp/auth.py @@ -713,6 +713,7 @@ async def email_auth_verify_route(request: web.Request) -> web.Response: session, referral_param, current_user_id=None, + settings=settings, ) db_user, _ = await user_dal.create_email_user( session, @@ -821,6 +822,7 @@ async def email_auth_magic_route(request: web.Request) -> web.Response: session, referral_param, current_user_id=None, + settings=settings, ) db_user, _ = await user_dal.create_email_user( session, @@ -1292,17 +1294,35 @@ async def _link_telegram_to_user( return current_user -def _normalize_referral_param(raw: Optional[str]) -> Optional[str]: +def _remnashop_referral_compat_enabled(settings: Optional[Settings]) -> bool: + if settings is None: + return False + return bool(getattr(settings, "MIGRATION_REMNASHOP_REFERRAL_CODE_COMPAT_ENABLED", False)) + + +def _strip_referral_param_prefix( + raw: Optional[str], + *, + preserve_current_u_prefix: bool, +) -> str: value = (raw or "").strip() if not value: - return None + return "" value_lower = value.lower() - if value_lower.startswith("ref_u"): + if value_lower.startswith("ref_u") and not preserve_current_u_prefix: value = value[5:] elif value_lower.startswith("ref_"): value = value[4:] - elif value and value[0].lower() == "u" and len(value) == 10: + return value + + +def _normalize_referral_param(raw: Optional[str]) -> Optional[str]: + value = _strip_referral_param_prefix(raw, preserve_current_u_prefix=False) + if not value: + return None + + if value and value[0].lower() == "u" and len(value) == 10: value = value[1:] if not re.fullmatch(r"[A-Za-z0-9]{1,32}", value): @@ -1310,26 +1330,64 @@ def _normalize_referral_param(raw: Optional[str]) -> Optional[str]: return value.upper() +def _referral_param_lookup_candidates( + raw: Optional[str], + *, + remnashop_compat: bool, +) -> List[str]: + if not remnashop_compat: + normalized = _normalize_referral_param(raw) + return [normalized] if normalized else [] + + value = _strip_referral_param_prefix(raw, preserve_current_u_prefix=True) + if not value or not re.fullmatch(r"[A-Za-z0-9._:-]{1,128}", value): + return [] + + candidates = [value] + if value and value[0].lower() == "u": + candidates.append(value[1:]) + + unique: List[str] = [] + for candidate in candidates: + if candidate and candidate not in unique: + unique.append(candidate) + return unique + + async def _resolve_referrer_id( session: AsyncSession, raw_referral_param: Optional[str], *, current_user_id: Optional[int], + settings: Optional[Settings] = None, ) -> Optional[int]: - normalized = _normalize_referral_param(raw_referral_param) - if not normalized: + remnashop_compat = _remnashop_referral_compat_enabled(settings) + candidates = _referral_param_lookup_candidates( + raw_referral_param, + remnashop_compat=remnashop_compat, + ) + if not candidates: return None - ref_user = None - if normalized.isdigit(): - ref_user = await user_dal.get_user_by_id(session, int(normalized)) - if not ref_user: - ref_user = await user_dal.get_user_by_referral_code(session, normalized) - if not ref_user: - return None - if current_user_id is not None and int(ref_user.user_id) == int(current_user_id): - return None - return int(ref_user.user_id) + for normalized in candidates: + ref_user = None + if normalized.isdigit() and not remnashop_compat: + ref_user = await user_dal.get_user_by_id(session, int(normalized)) + if not ref_user: + ref_user = await user_dal.get_user_by_referral_code( + session, + normalized, + include_legacy=remnashop_compat, + ) + if not ref_user and normalized.isdigit() and remnashop_compat: + ref_user = await user_dal.get_user_by_id(session, int(normalized)) + if not ref_user: + continue + if current_user_id is not None and int(ref_user.user_id) == int(current_user_id): + continue + return int(ref_user.user_id) + + return None async def _apply_referral_to_existing_user( @@ -1345,6 +1403,7 @@ async def _apply_referral_to_existing_user( session, raw_referral_param, current_user_id=int(user.user_id), + settings=request.app["settings"], ) if not referred_by_id: return False @@ -1427,6 +1486,7 @@ async def _ensure_user_from_telegram( session, referral_param or telegram_user.get("start_param"), current_user_id=user_id, + settings=settings, ) db_user, created = await user_dal.create_user( session, diff --git a/backend/bot/handlers/user/start.py b/backend/bot/handlers/user/start.py index 44cc598..79425a8 100644 --- a/backend/bot/handlers/user/start.py +++ b/backend/bot/handlers/user/start.py @@ -40,6 +40,67 @@ from db.models import User router = Router(name="user_start_router") +def _remnashop_referral_compat_enabled(settings: Settings) -> bool: + return bool(getattr(settings, "MIGRATION_REMNASHOP_REFERRAL_CODE_COMPAT_ENABLED", False)) + + +def _referral_code_lookup_candidates( + raw_ref_value: str, + *, + remnashop_compat: bool, +) -> list[str]: + value = str(raw_ref_value or "").strip() + if not value: + return [] + + candidates = [value] + if value and value[0].lower() == "u": + stripped_current_prefix = value[1:] + if remnashop_compat: + candidates.append(stripped_current_prefix) + else: + candidates = [stripped_current_prefix] + + unique: list[str] = [] + for candidate in candidates: + candidate = candidate.strip() + if candidate and candidate not in unique: + unique.append(candidate) + return unique + + +async def _resolve_referrer_from_start_ref( + session: AsyncSession, + raw_ref_value: str, + *, + settings: Settings, + current_user_id: int, +) -> Optional[int]: + ref_user: Optional[User] = None + if raw_ref_value.isdigit() and settings.LEGACY_REFS: + potential_referrer_id = int(raw_ref_value) + if potential_referrer_id != current_user_id: + ref_user = await user_dal.get_user_by_id(session, potential_referrer_id) + + include_legacy = _remnashop_referral_compat_enabled(settings) + if not ref_user: + for code in _referral_code_lookup_candidates( + raw_ref_value, + remnashop_compat=include_legacy, + ): + ref_user = await user_dal.get_user_by_referral_code( + session, + code, + include_legacy=include_legacy, + ) + if ref_user: + break + + if ref_user and ref_user.user_id != current_user_id: + return int(ref_user.user_id) + return None + + async def should_show_trial_button( settings: Settings, subscription_service: SubscriptionService, @@ -412,12 +473,10 @@ async def ensure_required_channel_subscription( @router.message(CommandStart()) @router.message( CommandStart( - magic=F.args.regexp(r"^ref_((?:[uU][A-Za-z0-9]{9})|(?:[A-Za-z0-9]{9})|\d+)$").as_( - "ref_match" - ) + magic=F.args.regexp(r"^ref_([A-Za-z0-9_-]{1,64})$").as_("ref_match") ) ) -@router.message(CommandStart(magic=F.args.regexp(r"^promo_(\w+)$").as_("promo_match"))) +@router.message(CommandStart(magic=F.args.regexp(r"^promo_([A-Za-z0-9_-]{1,100})$").as_("promo_match"))) @router.message(CommandStart(magic=F.args.regexp(r"^admin_user_(\d+)$").as_("admin_user_match"))) @router.message(CommandStart(magic=F.args.regexp(r"^ticket_(\d+)$").as_("ticket_match"))) @router.message(CommandStart(magic=F.args.regexp(r"^notifications$").as_("notifications_match"))) @@ -534,22 +593,12 @@ async def start_command_handler( if ref_match: raw_ref_value = ref_match.group(1) - if raw_ref_value.isdigit(): - if settings.LEGACY_REFS: - potential_referrer_id = int(raw_ref_value) - if potential_referrer_id != user_id and await user_dal.get_user_by_id( - session, potential_referrer_id - ): - referred_by_user_id = potential_referrer_id - else: - normalized_code = raw_ref_value.strip() - if normalized_code and normalized_code[0].lower() == "u": - normalized_code = normalized_code[1:] - ref_user = None - if normalized_code: - ref_user = await user_dal.get_user_by_referral_code(session, normalized_code) - if ref_user and ref_user.user_id != user_id: - referred_by_user_id = ref_user.user_id + referred_by_user_id = await _resolve_referrer_from_start_ref( + session, + raw_ref_value, + settings=settings, + current_user_id=user_id, + ) elif promo_match: promo_code_to_apply = promo_match.group(1) logging.info(f"User {user_id} started with promo code: {promo_code_to_apply}") diff --git a/backend/bot/services/promo_code_service.py b/backend/bot/services/promo_code_service.py index b486595..6f6a63d 100644 --- a/backend/bot/services/promo_code_service.py +++ b/backend/bot/services/promo_code_service.py @@ -38,8 +38,12 @@ class PromoCodeService: user_lang: str, ) -> Tuple[bool, datetime | str]: _ = lambda k, **kw: self.i18n.gettext(user_lang, k, **kw) - code_input_upper = (code_input or "").strip().upper()[:100] - code_display = html_escape(code_input_upper[:100], quote=False) + preserve_case = bool( + getattr(self.settings, "MIGRATION_REMNASHOP_PROMO_CODE_COMPAT_ENABLED", False) + ) + code_input_clean = (code_input or "").strip()[:100] + lookup_code = code_input_clean if preserve_case else code_input_clean.upper() + code_display = html_escape(lookup_code[:100], quote=False) throttle_identifier = self._throttle_identifier(user_id) throttle = await security_dal.check_throttle( @@ -54,7 +58,7 @@ class PromoCodeService: ) promo_data = await promo_code_dal.get_active_promo_code_by_code_str( - session, code_input_upper + session, lookup_code, preserve_case=preserve_case ) if not promo_data: @@ -71,9 +75,11 @@ class PromoCodeService: "promo_code_too_many_attempts", seconds=throttle_result.retry_after or max(1, int(self.settings.BRUTE_FORCE_LOCK_SECONDS)), - ) + ) return False, _("promo_code_not_found", code=code_display) + applied_code = str(promo_data.code or lookup_code) + code_display = html_escape(applied_code[:100], quote=False) existing_activation = await promo_code_dal.get_user_activation_for_promo( session, promo_data.promo_code_id, user_id ) @@ -86,7 +92,7 @@ class PromoCodeService: session=session, user_id=user_id, bonus_days=bonus_days, - reason=f"promo code {code_input_upper}", + reason=f"promo code {applied_code}", ) if new_end_date: @@ -109,7 +115,7 @@ class PromoCodeService: user = await user_dal.get_user_by_id(session, user_id) await notification_service.notify_promo_activation( user_id=user_id, - promo_code=code_input_upper, + promo_code=applied_code, bonus_days=bonus_days, username=user.username if user else None, email=getattr(user, "email", None) if user else None, diff --git a/backend/config/settings.py b/backend/config/settings.py index bd09406..3c856c7 100644 --- a/backend/config/settings.py +++ b/backend/config/settings.py @@ -286,6 +286,24 @@ class Settings(BaseSettings): default=True, description="Allow legacy referral links like ref_ to continue working. Defaults to True when unset.", # noqa: E501 ) + MIGRATION_REMNASHOP_REFERRAL_CODE_COMPAT_ENABLED: bool = Field( + default=False, + description=( + "Accept referral links imported from snoups/remnashop via legacy_referral_codes." + ), + ) + MIGRATION_REMNASHOP_PROMO_CODE_COMPAT_ENABLED: bool = Field( + default=False, + description="Try exact legacy Remnashop promo codes before uppercase normalization.", + ) + MIGRATION_REMNASHOP_IMPORTED_AT: Optional[str] = Field( + default=None, + description="Timestamp of the latest Remnashop import run, managed by the import script.", + ) + MIGRATION_REMNASHOP_NOTES: Optional[str] = Field( + default=None, + description="Operator notes for instances migrated from Remnashop.", + ) APP_RUNTIME_MODE: str = Field( default="production", diff --git a/backend/db/dal/promo_code_dal.py b/backend/db/dal/promo_code_dal.py index 8047792..41a376b 100644 --- a/backend/db/dal/promo_code_dal.py +++ b/backend/db/dal/promo_code_dal.py @@ -22,24 +22,46 @@ async def get_promo_code_by_id(session: AsyncSession, promo_code_id: int) -> Opt return await session.get(PromoCode, promo_code_id) -async def get_promo_code_by_code(session: AsyncSession, code_str: str) -> Optional[PromoCode]: +def _promo_lookup_candidates(code_str: str, *, preserve_case: bool) -> List[str]: + code = str(code_str or "").strip() + if not code: + return [] + candidates = [code] if preserve_case else [] + upper_code = code.upper() + if upper_code not in candidates: + candidates.append(upper_code) + return candidates + + +async def get_promo_code_by_code( + session: AsyncSession, code_str: str, *, preserve_case: bool = False +) -> Optional[PromoCode]: """Get promo code by code string (regardless of active status)""" - stmt = select(PromoCode).where(PromoCode.code == code_str.upper()) - result = await session.execute(stmt) - return result.scalar_one_or_none() + for candidate in _promo_lookup_candidates(code_str, preserve_case=preserve_case): + stmt = select(PromoCode).where(PromoCode.code == candidate) + result = await session.execute(stmt) + promo = result.scalar_one_or_none() + if promo: + return promo + return None async def get_active_promo_code_by_code_str( - session: AsyncSession, code_str: str + session: AsyncSession, code_str: str, *, preserve_case: bool = False ) -> Optional[PromoCode]: - stmt = select(PromoCode).where( - PromoCode.code == code_str.upper(), - PromoCode.is_active == True, - PromoCode.current_activations < PromoCode.max_activations, - or_(PromoCode.valid_until == None, PromoCode.valid_until > datetime.now(timezone.utc)), - ) - result = await session.execute(stmt) - return result.scalar_one_or_none() + now = datetime.now(timezone.utc) + for candidate in _promo_lookup_candidates(code_str, preserve_case=preserve_case): + stmt = select(PromoCode).where( + PromoCode.code == candidate, + PromoCode.is_active == True, + PromoCode.current_activations < PromoCode.max_activations, + or_(PromoCode.valid_until == None, PromoCode.valid_until > now), + ) + result = await session.execute(stmt) + promo = result.scalar_one_or_none() + if promo: + return promo + return None async def get_all_active_promo_codes( diff --git a/backend/db/dal/user_dal.py b/backend/db/dal/user_dal.py index bc58eef..0acc533 100644 --- a/backend/db/dal/user_dal.py +++ b/backend/db/dal/user_dal.py @@ -4,7 +4,7 @@ import string from datetime import datetime, timedelta, timezone from typing import Any, Dict, List, Optional, Tuple -from sqlalchemy import and_, case, delete, desc, func, or_, update +from sqlalchemy import String, and_, case, cast, delete, desc, func, or_, update from sqlalchemy.dialects.postgresql import insert as pg_insert from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.future import select @@ -14,6 +14,8 @@ from ..models import ( AdAttribution, EmailVerificationCode, HwidDevicePurchase, + LegacyImportMapping, + LegacyReferralCode, MessageLog, Payment, PromoCodeActivation, @@ -76,7 +78,7 @@ async def ensure_referral_code(session: AsyncSession, user: User) -> str: Returns the existing or newly generated code. """ if user.referral_code: - normalized = user.referral_code.strip().upper() + normalized = user.referral_code.strip() if normalized != user.referral_code: user.referral_code = normalized await session.flush() @@ -210,7 +212,7 @@ async def create_user(session: AsyncSession, user_data: Dict[str, Any]) -> Tuple if not user_data.get("referral_code"): user_data["referral_code"] = await generate_unique_referral_code(session) else: - user_data["referral_code"] = user_data["referral_code"].strip().upper() + user_data["referral_code"] = user_data["referral_code"].strip() # Use PostgreSQL upsert to avoid IntegrityError on concurrent inserts stmt = ( @@ -567,6 +569,19 @@ async def merge_users( await session.execute( update(model).where(model.user_id == source_user_id).values(user_id=target_user_id) ) + await session.execute( + update(LegacyReferralCode) + .where(LegacyReferralCode.user_id == source_user_id) + .values(user_id=target_user_id) + ) + await session.execute( + update(LegacyImportMapping) + .where( + LegacyImportMapping.target_table == "users", + LegacyImportMapping.target_id == str(source_user_id), + ) + .values(target_id=str(target_user_id)) + ) await session.execute( update(MessageLog) @@ -590,13 +605,60 @@ async def merge_users( return target -async def get_user_by_referral_code(session: AsyncSession, referral_code: str) -> Optional[User]: - normalized = referral_code.strip().upper() +async def get_user_by_referral_code( + session: AsyncSession, + referral_code: str, + *, + include_legacy: bool = False, +) -> Optional[User]: + normalized = referral_code.strip() if not normalized: return None + stmt = select(User).where(User.referral_code == normalized) result = await session.execute(stmt) - return result.scalar_one_or_none() + user = result.scalar_one_or_none() + if user: + return user + + upper_normalized = normalized.upper() + if upper_normalized != normalized: + stmt = select(User).where(User.referral_code == upper_normalized) + result = await session.execute(stmt) + user = result.scalar_one_or_none() + if user: + return user + + if not include_legacy: + return None + + stmt = ( + select(User) + .join(LegacyReferralCode, LegacyReferralCode.user_id == User.user_id) + .where(LegacyReferralCode.code == normalized, LegacyReferralCode.is_active == True) + .limit(1) + ) + result = await session.execute(stmt) + user = result.scalar_one_or_none() + if user: + return user + + if upper_normalized != normalized: + stmt = ( + select(User) + .join(LegacyReferralCode, LegacyReferralCode.user_id == User.user_id) + .where( + LegacyReferralCode.code == upper_normalized, + LegacyReferralCode.is_active == True, + ) + .limit(1) + ) + result = await session.execute(stmt) + user = result.scalar_one_or_none() + if user: + return user + + return None async def update_user( @@ -1020,6 +1082,31 @@ async def delete_user_and_relations(session: AsyncSession, user_id: int) -> bool await session.execute(delete(UserBilling).where(UserBilling.user_id == user_id)) await session.execute(delete(AdAttribution).where(AdAttribution.user_id == user_id)) await session.execute(delete(UserTelegramAvatar).where(UserTelegramAvatar.user_id == user_id)) + await session.execute(delete(LegacyReferralCode).where(LegacyReferralCode.user_id == user_id)) + await session.execute( + delete(LegacyImportMapping).where( + or_( + and_( + LegacyImportMapping.target_table == "users", + LegacyImportMapping.target_id == str(user_id), + ), + and_( + LegacyImportMapping.target_table == "subscriptions", + LegacyImportMapping.target_id.in_( + select(cast(Subscription.subscription_id, String)).where( + Subscription.user_id == user_id + ) + ), + ), + and_( + LegacyImportMapping.target_table == "payments", + LegacyImportMapping.target_id.in_( + select(cast(Payment.payment_id, String)).where(Payment.user_id == user_id) + ), + ), + ) + ) + ) await session.execute(delete(Payment).where(Payment.user_id == user_id)) await session.execute(delete(Subscription).where(Subscription.user_id == user_id)) diff --git a/backend/db/migrator.py b/backend/db/migrator.py index 73b3e3c..18a76fa 100644 --- a/backend/db/migrator.py +++ b/backend/db/migrator.py @@ -1070,6 +1070,79 @@ def _migration_0033_add_trial_eligibility_reset_marker(connection: Connection) - ) +def _migration_0034_add_legacy_import_compatibility(connection: Connection) -> None: + inspector = inspect(connection) + table_names = set(inspector.get_table_names()) + + if "users" in table_names: + columns = {col["name"]: col for col in inspector.get_columns("users")} + referral_column = columns.get("referral_code") + length = getattr(referral_column.get("type"), "length", None) if referral_column else None + if referral_column and (length is None or int(length) < 64): + connection.execute( + text("ALTER TABLE users ALTER COLUMN referral_code TYPE VARCHAR(64)") + ) + + connection.execute( + text( + """ + CREATE TABLE IF NOT EXISTS legacy_referral_codes ( + legacy_code_id SERIAL PRIMARY KEY, + source VARCHAR(64) NOT NULL DEFAULT 'remnashop', + code VARCHAR(128) NOT NULL, + user_id BIGINT NOT NULL REFERENCES users(user_id), + is_active BOOLEAN NOT NULL DEFAULT TRUE, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NULL, + CONSTRAINT uq_legacy_referral_source_code UNIQUE (source, code) + ) + """ + ) + ) + for stmt in [ + ( + "CREATE INDEX IF NOT EXISTS ix_legacy_referral_codes_source " + "ON legacy_referral_codes (source)" + ), + "CREATE INDEX IF NOT EXISTS ix_legacy_referral_codes_code ON legacy_referral_codes (code)", + ( + "CREATE INDEX IF NOT EXISTS ix_legacy_referral_codes_user_id " + "ON legacy_referral_codes (user_id)" + ), + ( + "CREATE INDEX IF NOT EXISTS ix_legacy_referral_codes_is_active " + "ON legacy_referral_codes (is_active)" + ), + ]: + connection.execute(text(stmt)) + + connection.execute( + text( + """ + CREATE TABLE IF NOT EXISTS legacy_import_mappings ( + source VARCHAR(64) NOT NULL, + entity_type VARCHAR(64) NOT NULL, + source_id VARCHAR(128) NOT NULL, + target_table VARCHAR(128) NOT NULL, + target_id VARCHAR(128) NOT NULL, + metadata_json TEXT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NULL, + PRIMARY KEY (source, entity_type, source_id) + ) + """ + ) + ) + connection.execute( + text( + """ + CREATE INDEX IF NOT EXISTS ix_legacy_import_mappings_target + ON legacy_import_mappings (target_table, target_id) + """ + ) + ) + + MIGRATIONS: List[Migration] = [ Migration( id="0001_add_channel_subscription_fields", @@ -1247,6 +1320,11 @@ MIGRATIONS: List[Migration] = [ description="Track admin resets of per-user trial eligibility without deleting history", upgrade=_migration_0033_add_trial_eligibility_reset_marker, ), + Migration( + id="0034_add_legacy_import_compatibility", + description="Store legacy import mappings and referral codes for source-bot migrations", + upgrade=_migration_0034_add_legacy_import_compatibility, + ), ] diff --git a/backend/db/models.py b/backend/db/models.py index 78bb8c3..e524d0d 100644 --- a/backend/db/models.py +++ b/backend/db/models.py @@ -43,7 +43,7 @@ class User(Base): registration_date = Column(DateTime(timezone=True), server_default=func.now()) is_banned = Column(Boolean, default=False) panel_user_uuid = Column(String, nullable=True, unique=True, index=True) - referral_code = Column(String(16), nullable=True, unique=True, index=True) + referral_code = Column(String(64), nullable=True, unique=True, index=True) referred_by_id = Column(BigInteger, ForeignKey("users.user_id"), nullable=True) lifetime_used_traffic_bytes = Column(BigInteger, nullable=True) lifetime_used_traffic_synced_at = Column(DateTime(timezone=True), nullable=True) @@ -396,6 +396,37 @@ class PromoCodeActivation(Base): ) +class LegacyReferralCode(Base): + __tablename__ = "legacy_referral_codes" + + legacy_code_id = Column(Integer, primary_key=True, autoincrement=True) + source = Column(String(64), nullable=False, default="remnashop", index=True) + code = Column(String(128), nullable=False, index=True) + user_id = Column(BigInteger, ForeignKey("users.user_id"), nullable=False, index=True) + is_active = Column(Boolean, nullable=False, default=True, index=True) + created_at = Column(DateTime(timezone=True), server_default=func.now()) + updated_at = Column(DateTime(timezone=True), onupdate=func.now(), nullable=True) + + user = relationship("User") + + __table_args__ = ( + UniqueConstraint("source", "code", name="uq_legacy_referral_source_code"), + ) + + +class LegacyImportMapping(Base): + __tablename__ = "legacy_import_mappings" + + source = Column(String(64), primary_key=True) + entity_type = Column(String(64), primary_key=True) + source_id = Column(String(128), primary_key=True) + target_table = Column(String(128), nullable=False) + target_id = Column(String(128), nullable=False) + metadata_json = Column(Text, nullable=True) + created_at = Column(DateTime(timezone=True), server_default=func.now()) + updated_at = Column(DateTime(timezone=True), onupdate=func.now(), nullable=True) + + class MessageLog(Base): __tablename__ = "message_logs" diff --git a/backend/scripts/__init__.py b/backend/scripts/__init__.py new file mode 100644 index 0000000..7e7e6c0 --- /dev/null +++ b/backend/scripts/__init__.py @@ -0,0 +1,2 @@ +"""Operational one-shot scripts shipped with the backend image.""" + diff --git a/backend/scripts/import_legacy.py b/backend/scripts/import_legacy.py new file mode 100644 index 0000000..b2fede7 --- /dev/null +++ b/backend/scripts/import_legacy.py @@ -0,0 +1,1164 @@ +"""Import data from legacy source bots into the current shop database. + +Currently supported source: + remnashop + +Example: + python backend/scripts/import_legacy.py \ + --source-type remnashop \ + --source-dsn postgresql://user:pass@localhost:5432/remnashop \ + --dry-run +""" + +from __future__ import annotations + +import argparse +import asyncio +import json +import logging +import re +import sys +from collections import defaultdict +from datetime import datetime, timedelta, timezone +from decimal import Decimal, InvalidOperation +from pathlib import Path +from typing import Any, Iterable, Optional + +from sqlalchemy import inspect, select, text +from sqlalchemy.dialects.postgresql import insert as pg_insert +from sqlalchemy.ext.asyncio import ( + AsyncConnection, + AsyncSession, + async_sessionmaker, + create_async_engine, +) + +BACKEND_ROOT = Path(__file__).resolve().parents[1] +if str(BACKEND_ROOT) not in sys.path: + sys.path.insert(0, str(BACKEND_ROOT)) + +from config.settings import Settings # noqa: E402 +from db.dal import user_dal # noqa: E402 +from db.migrator import run_database_migrations # noqa: E402 +from db.models import ( # noqa: E402 + AppSettingOverride, + Base, + LegacyImportMapping, + LegacyReferralCode, + MessageLog, + Payment, + PromoCode, + PromoCodeActivation, + Subscription, + User, +) + +SOURCE = "remnashop" +GIB = 1024**3 +UUID_RE = re.compile( + r"\b[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-" + r"[0-9a-fA-F]{4}-[0-9a-fA-F]{12}\b" +) +SAFE_SCHEMA_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") + +logger = logging.getLogger(__name__) + + +def normalize_async_postgres_dsn(dsn: str) -> str: + value = str(dsn or "").strip() + if value.startswith("postgresql+asyncpg://"): + return value + if value.startswith("postgresql://"): + return "postgresql+asyncpg://" + value.removeprefix("postgresql://") + if value.startswith("postgres://"): + return "postgresql+asyncpg://" + value.removeprefix("postgres://") + return value + + +def _json_default(value: Any) -> str: + if isinstance(value, (datetime, Decimal)): + return str(value) + return str(value) + + +def _json_dumps(value: Any) -> str: + return json.dumps(value, ensure_ascii=False, sort_keys=True, default=_json_default) + + +def _safe_schema_name(schema: str) -> str: + value = str(schema or "public").strip() + if not SAFE_SCHEMA_RE.fullmatch(value): + raise ValueError(f"Unsafe PostgreSQL schema name: {schema!r}") + return value + + +def _qtable(schema: str, table: str) -> str: + schema = _safe_schema_name(schema) + return f'"{schema}"."{table}"' + + +def _as_mapping(row: Any) -> dict[str, Any]: + return dict(row._mapping if hasattr(row, "_mapping") else row) + + +def _as_utc(value: Any) -> Optional[datetime]: + if value is None: + return None + if isinstance(value, datetime): + result = value + else: + text_value = str(value).strip() + if not text_value: + return None + try: + result = datetime.fromisoformat(text_value.replace("Z", "+00:00")) + except ValueError: + return None + if result.tzinfo is None: + return result.replace(tzinfo=timezone.utc) + return result.astimezone(timezone.utc) + + +def _to_decimal(value: Any) -> Optional[Decimal]: + if value is None: + return None + try: + return Decimal(str(value)) + except (InvalidOperation, ValueError): + return None + + +def _to_int(value: Any) -> Optional[int]: + number = _to_decimal(value) + if number is None: + return None + try: + return int(number) + except (OverflowError, ValueError): + return None + + +def _split_name(name: Any) -> tuple[Optional[str], Optional[str]]: + value = str(name or "").strip() + if not value: + return None, None + parts = value.split(maxsplit=1) + if len(parts) == 1: + return parts[0][:255], None + return parts[0][:255], parts[1][:255] + + +def _jsonish(value: Any) -> dict[str, Any]: + if isinstance(value, dict): + return value + if isinstance(value, str) and value.strip(): + try: + decoded = json.loads(value) + except ValueError: + return {} + return decoded if isinstance(decoded, dict) else {} + return {} + + +def _listish(value: Any) -> list[Any]: + if value is None: + return [] + if isinstance(value, list): + return value + if isinstance(value, tuple): + return list(value) + return [value] + + +def remnashop_traffic_gb_to_bytes(value: Any) -> Optional[int]: + number = _to_decimal(value) + if number is None: + return None + return int(number * GIB) + + +def remnashop_pricing_amount(pricing: Any) -> float: + data = _jsonish(pricing) + for key in ("final_amount", "total_amount", "amount", "price"): + number = _to_decimal(data.get(key)) + if number is not None: + return float(number) + return 0.0 + + +def remnashop_pricing_currency(pricing: Any, fallback: Any = None) -> str: + data = _jsonish(pricing) + currency = str(data.get("currency") or fallback or "RUB").strip().upper() + return currency or "RUB" + + +def remnashop_transaction_status(status: Any, gateway_type: Any = None) -> str: + source_status = str(status or "").strip().upper() + provider = str(gateway_type or "").strip().lower() + if source_status == "COMPLETED": + return "succeeded" + if source_status == "PENDING": + return f"pending_{provider}" if provider else "pending" + if source_status == "CANCELED": + return "canceled" + if source_status == "REFUNDED": + return "refunded" + if source_status == "FAILED": + return "failed" + return source_status.lower() or "unknown" + + +def remnashop_sale_mode(purchase_type: Any) -> str: + source_type = str(purchase_type or "").strip().upper() + if source_type in {"NEW", "RENEW"}: + return "subscription" + if source_type == "CHANGE": + return "tariff_upgrade" + return source_type.lower() or "subscription" + + +def remnashop_months_from_plan_snapshot( + plan_snapshot: Any, + *, + created_at: Any = None, + expire_at: Any = None, +) -> Optional[int]: + data = _jsonish(plan_snapshot) + for key in ("duration_months", "months", "month"): + months = _to_int(data.get(key)) + if months and months > 0: + return months + + for key in ("duration_days", "days", "duration"): + days = _to_int(data.get(key)) + if days and days > 0: + return max(1, round(days / 30)) + + start = _as_utc(created_at) + end = _as_utc(expire_at) + if start and end and end > start: + return max(1, round((end - start).days / 30)) + return None + + +def remnashop_tariff_key(plan_snapshot: Any, tariff_map: dict[str, str]) -> Optional[str]: + data = _jsonish(plan_snapshot) + candidates = [ + data.get("id"), + data.get("name"), + data.get("tag"), + data.get("public_code"), + ] + for candidate in candidates: + key = str(candidate or "").strip() + if key and key in tariff_map: + return tariff_map[key] + return None + + +def _provider_value(gateway_type: Any) -> str: + value = str(gateway_type or "remnashop").strip().lower() + if value == "telegram_stars": + return "stars" + return value or "remnashop" + + +def _extract_panel_subscription_uuid(url: Any, panel_user_uuid: Optional[str]) -> Optional[str]: + value = str(url or "") + if not value: + return None + panel_user_uuid = str(panel_user_uuid or "").lower() + for match in UUID_RE.finditer(value): + candidate = match.group(0).lower() + if candidate != panel_user_uuid: + return candidate + return None + + +def _legacy_user_metadata(row: dict[str, Any]) -> dict[str, Any]: + keys = ( + "id", + "points", + "personal_discount", + "purchase_discount", + "role", + "is_rules_accepted", + "is_trial_available", + "language", + "current_subscription_id", + ) + return {key: row.get(key) for key in keys if row.get(key) is not None} + + +def _counter() -> dict[str, int]: + return defaultdict(int) + + +class RemnashopImporter: + def __init__( + self, + *, + source: AsyncConnection, + target: AsyncSession, + source_schema: str, + only: set[str], + on_conflict: str, + dry_run: bool, + created_by_admin_id: int, + tariff_map: dict[str, str], + write_admin_compat_overrides: bool, + ) -> None: + self.source = source + self.target = target + self.source_schema = _safe_schema_name(source_schema) + self.only = only + self.on_conflict = on_conflict + self.dry_run = dry_run + self.created_by_admin_id = created_by_admin_id + self.tariff_map = tariff_map + self.write_admin_compat_overrides = write_admin_compat_overrides + self.tables: set[str] = set() + self.user_map: dict[int, int] = {} + self.summary: dict[str, Any] = { + "source": SOURCE, + "dry_run": dry_run, + "on_conflict": on_conflict, + "users": _counter(), + "referrals": _counter(), + "subscriptions": _counter(), + "payments": _counter(), + "promocodes": _counter(), + "settings": _counter(), + "warnings": [], + } + + async def run(self) -> dict[str, Any]: + self.tables = await self._source_tables() + await self._warn_missing_tables() + + if self._should_run("users"): + await self.import_users() + if self._should_run("referrals"): + await self.import_referrals() + if self._should_run("subscriptions"): + await self.import_subscriptions() + if self._should_run("payments"): + await self.import_payments() + if self._should_run("promocodes"): + await self.import_promocodes() + if self._should_run("settings"): + await self.import_settings() + + if self.write_admin_compat_overrides: + await self._write_admin_overrides() + + return self._plain_summary() + + def _plain_summary(self) -> dict[str, Any]: + result = dict(self.summary) + for key, value in list(result.items()): + if isinstance(value, defaultdict): + result[key] = dict(value) + return result + + def _should_run(self, key: str) -> bool: + return not self.only or key in self.only or "all" in self.only + + async def _source_tables(self) -> set[str]: + def load_tables(sync_connection: Any) -> set[str]: + return set(inspect(sync_connection).get_table_names(schema=self.source_schema)) + + return await self.source.run_sync(load_tables) + + async def _warn_missing_tables(self) -> None: + required = {"users", "subscriptions", "transactions", "referrals", "settings"} + missing = sorted(required - self.tables) + if missing: + self.summary["warnings"].append(f"Missing source tables: {', '.join(missing)}") + + async def _fetch_rows(self, table: str, *, order_by: str = "id") -> list[dict[str, Any]]: + if table not in self.tables: + return [] + order_sql = f" ORDER BY {order_by}" if order_by else "" + result = await self.source.execute( + text(f"SELECT * FROM {_qtable(self.source_schema, table)}{order_sql}") + ) + return [_as_mapping(row) for row in result.mappings().all()] + + async def _fetch_one(self, table: str) -> Optional[dict[str, Any]]: + rows = await self._fetch_rows(table, order_by="") + return rows[0] if rows else None + + async def _latest_panel_uuid_by_telegram(self) -> dict[int, str]: + if "subscriptions" not in self.tables: + return {} + result = await self.source.execute( + text( + f""" + SELECT DISTINCT ON (user_telegram_id) + user_telegram_id, + user_remna_id + FROM {_qtable(self.source_schema, "subscriptions")} + WHERE user_remna_id IS NOT NULL + ORDER BY user_telegram_id, updated_at DESC NULLS LAST, id DESC + """ + ) + ) + panel_by_tg: dict[int, str] = {} + for row in result.mappings().all(): + telegram_id = _to_int(row.get("user_telegram_id")) + panel_uuid = str(row.get("user_remna_id") or "").strip() + if telegram_id and panel_uuid: + panel_by_tg[telegram_id] = panel_uuid + return panel_by_tg + + async def _target_user_for_telegram(self, telegram_id: Any) -> Optional[User]: + normalized = _to_int(telegram_id) + if normalized is None: + return None + user = await user_dal.get_user_by_telegram_id(self.target, normalized) + if not user: + user = await user_dal.get_user_by_id(self.target, normalized) + if user: + self.user_map[normalized] = int(user.user_id) + return user + + def _can_overwrite(self) -> bool: + return self.on_conflict == "overwrite" + + def _can_merge_existing(self) -> bool: + return self.on_conflict in {"merge", "overwrite"} + + def _assign_if_allowed(self, model: Any, attr: str, value: Any) -> bool: + if value is None: + return False + current = getattr(model, attr, None) + if self._can_overwrite() or current in (None, ""): + setattr(model, attr, value) + return True + return False + + async def _upsert_mapping( + self, + *, + entity_type: str, + source_id: Any, + target_table: str, + target_id: Any, + metadata: Optional[dict[str, Any]] = None, + ) -> None: + now = datetime.now(timezone.utc) + source_id_value = str(source_id) + target_id_value = str(target_id) + stmt = ( + pg_insert(LegacyImportMapping) + .values( + source=SOURCE, + entity_type=entity_type, + source_id=source_id_value, + target_table=target_table, + target_id=target_id_value, + metadata_json=_json_dumps(metadata or {}), + updated_at=now, + ) + .on_conflict_do_update( + index_elements=[ + LegacyImportMapping.source, + LegacyImportMapping.entity_type, + LegacyImportMapping.source_id, + ], + set_={ + "target_table": target_table, + "target_id": target_id_value, + "metadata_json": _json_dumps(metadata or {}), + "updated_at": now, + }, + ) + ) + await self.target.execute(stmt) + + async def _get_mapping(self, entity_type: str, source_id: Any) -> Optional[LegacyImportMapping]: + stmt = select(LegacyImportMapping).where( + LegacyImportMapping.source == SOURCE, + LegacyImportMapping.entity_type == entity_type, + LegacyImportMapping.source_id == str(source_id), + ) + result = await self.target.execute(stmt) + return result.scalar_one_or_none() + + async def _upsert_setting_override(self, key: str, value: Any) -> None: + now = datetime.now(timezone.utc) + encoded = json.dumps(value, ensure_ascii=False, separators=(",", ":")) + stmt = ( + pg_insert(AppSettingOverride) + .values( + key=key, + value=encoded, + updated_at=now, + updated_by=self.created_by_admin_id or None, + ) + .on_conflict_do_update( + index_elements=[AppSettingOverride.key], + set_={ + "value": encoded, + "updated_at": now, + "updated_by": self.created_by_admin_id or None, + }, + ) + ) + await self.target.execute(stmt) + + async def _upsert_legacy_referral_code(self, *, code: str, user_id: int) -> None: + if len(code) > 128: + self.summary["warnings"].append( + f"Skipped overlong legacy referral code for user {user_id}: {len(code)} chars" + ) + return + now = datetime.now(timezone.utc) + stmt = ( + pg_insert(LegacyReferralCode) + .values( + source=SOURCE, + code=code, + user_id=user_id, + is_active=True, + updated_at=now, + ) + .on_conflict_do_update( + index_elements=[LegacyReferralCode.source, LegacyReferralCode.code], + set_={"user_id": user_id, "is_active": True, "updated_at": now}, + ) + ) + await self.target.execute(stmt) + + async def _record_user_state_note( + self, + *, + telegram_id: int, + user_id: int, + metadata: dict[str, Any], + ) -> None: + if not metadata: + return + if await self._get_mapping("user_state", telegram_id): + return + log = MessageLog( + user_id=None, + target_user_id=user_id, + event_type="legacy_remnashop_user_state", + content=_json_dumps(metadata), + is_admin_event=True, + ) + self.target.add(log) + await self.target.flush() + await self._upsert_mapping( + entity_type="user_state", + source_id=telegram_id, + target_table="message_logs", + target_id=log.log_id, + metadata=metadata, + ) + + async def _source_referral_code_conflicts(self, code: str, user_id: int) -> bool: + existing = await user_dal.get_user_by_referral_code( + self.target, + code, + include_legacy=False, + ) + return bool(existing and int(existing.user_id) != int(user_id)) + + async def import_users(self) -> None: + rows = await self._fetch_rows("users", order_by="telegram_id") + panel_by_tg = await self._latest_panel_uuid_by_telegram() + for row in rows: + telegram_id = _to_int(row.get("telegram_id")) + if telegram_id is None: + self.summary["users"]["skipped"] += 1 + continue + + first_name, last_name = _split_name(row.get("name")) + panel_uuid = panel_by_tg.get(telegram_id) + referral_code = str(row.get("referral_code") or "").strip() or None + created_at = _as_utc(row.get("created_at")) or datetime.now(timezone.utc) + language = str(row.get("language") or "ru").strip().lower()[:8] or "ru" + + existing = await self._target_user_for_telegram(telegram_id) + if existing and self.on_conflict == "skip": + target = existing + self.summary["users"]["skipped"] += 1 + elif existing: + target = existing + if self._can_merge_existing(): + self._assign_if_allowed(target, "username", row.get("username")) + self._assign_if_allowed(target, "first_name", first_name) + self._assign_if_allowed(target, "last_name", last_name) + self._assign_if_allowed(target, "language_code", language) + self._assign_if_allowed(target, "panel_user_uuid", panel_uuid) + if bool(row.get("is_blocked")): + target.is_banned = True + elif self._can_overwrite(): + target.is_banned = False + if bool(row.get("is_bot_blocked")): + target.telegram_notifications_status = "blocked" + target.telegram_notifications_checked_at = datetime.now(timezone.utc) + target.telegram_notifications_blocked_at = datetime.now(timezone.utc) + if referral_code and len(referral_code) <= 64 and not target.referral_code: + if not await self._source_referral_code_conflicts( + referral_code, + int(target.user_id), + ): + target.referral_code = referral_code + self.summary["users"]["updated"] += 1 + else: + new_referral_code = None + if referral_code and len(referral_code) <= 64: + conflict = await self._source_referral_code_conflicts( + referral_code, + telegram_id, + ) + if not conflict: + new_referral_code = referral_code + + target, created = await user_dal.create_user( + self.target, + { + "user_id": telegram_id, + "telegram_id": telegram_id, + "username": row.get("username"), + "first_name": first_name, + "last_name": last_name, + "language_code": language, + "registration_date": created_at, + "is_banned": bool(row.get("is_blocked")), + "panel_user_uuid": panel_uuid, + "referral_code": new_referral_code, + "telegram_notifications_status": "blocked" + if bool(row.get("is_bot_blocked")) + else "unknown", + "telegram_notifications_checked_at": datetime.now(timezone.utc) + if bool(row.get("is_bot_blocked")) + else None, + "telegram_notifications_blocked_at": datetime.now(timezone.utc) + if bool(row.get("is_bot_blocked")) + else None, + }, + ) + self.summary["users"]["created" if created else "updated"] += 1 + + if not target: + self.summary["users"]["skipped"] += 1 + continue + + self.user_map[telegram_id] = int(target.user_id) + if referral_code: + await self._upsert_legacy_referral_code(code=referral_code, user_id=target.user_id) + + metadata = _legacy_user_metadata(row) + if panel_uuid: + metadata["panel_user_uuid"] = panel_uuid + await self._upsert_mapping( + entity_type="user", + source_id=telegram_id, + target_table="users", + target_id=target.user_id, + metadata=metadata, + ) + await self._record_user_state_note( + telegram_id=telegram_id, + user_id=int(target.user_id), + metadata=metadata, + ) + + await self.target.flush() + + async def import_referrals(self) -> None: + rows = await self._fetch_rows("referrals", order_by="id") + for row in rows: + referrer = await self._target_user_for_telegram(row.get("referrer_telegram_id")) + referred = await self._target_user_for_telegram(row.get("referred_telegram_id")) + if not referrer or not referred or referrer.user_id == referred.user_id: + self.summary["referrals"]["skipped"] += 1 + continue + if referred.referred_by_id and not self._can_overwrite(): + self.summary["referrals"]["skipped"] += 1 + continue + referred.referred_by_id = int(referrer.user_id) + self.summary["referrals"]["updated"] += 1 + await self._upsert_mapping( + entity_type="referral", + source_id=row.get("id") or f"{referrer.user_id}:{referred.user_id}", + target_table="users", + target_id=referred.user_id, + metadata={ + "referrer_user_id": referrer.user_id, + "referred_user_id": referred.user_id, + }, + ) + await self.target.flush() + + async def import_subscriptions(self) -> None: + rows = await self._fetch_rows("subscriptions", order_by="id") + now = datetime.now(timezone.utc) + for row in rows: + user = await self._target_user_for_telegram(row.get("user_telegram_id")) + if not user: + self.summary["subscriptions"]["skipped"] += 1 + continue + + panel_user_uuid = str(row.get("user_remna_id") or user.panel_user_uuid or "").strip() + if not panel_user_uuid: + self.summary["subscriptions"]["skipped"] += 1 + continue + if not user.panel_user_uuid or self._can_overwrite(): + user.panel_user_uuid = panel_user_uuid + + source_id = row.get("id") + mapping = await self._get_mapping("subscription", source_id) + existing: Optional[Subscription] = None + if mapping and str(mapping.target_id).isdigit(): + existing = await self.target.get(Subscription, int(mapping.target_id)) + + panel_sub_uuid = _extract_panel_subscription_uuid(row.get("url"), panel_user_uuid) + if not existing and panel_sub_uuid: + existing = ( + await self.target.execute( + select(Subscription).where( + Subscription.panel_subscription_uuid == panel_sub_uuid + ) + ) + ).scalar_one_or_none() + + status = str(row.get("status") or "UNKNOWN").strip().upper() + expire_at = _as_utc(row.get("expire_at")) or now + created_at = _as_utc(row.get("created_at")) or now + plan_snapshot = _jsonish(row.get("plan_snapshot")) + traffic_limit_bytes = remnashop_traffic_gb_to_bytes(row.get("traffic_limit")) + payload = { + "user_id": int(user.user_id), + "panel_user_uuid": panel_user_uuid, + "panel_subscription_uuid": panel_sub_uuid, + "start_date": created_at, + "end_date": expire_at, + "duration_months": remnashop_months_from_plan_snapshot( + plan_snapshot, + created_at=created_at, + expire_at=expire_at, + ), + "is_active": status in {"ACTIVE", "LIMITED"} and expire_at > now, + "status_from_panel": status, + "traffic_limit_bytes": traffic_limit_bytes, + "provider": "trial" if bool(row.get("is_trial")) else SOURCE, + "skip_notifications": True, + "auto_renew_enabled": False, + "tariff_key": remnashop_tariff_key(plan_snapshot, self.tariff_map), + "tier_baseline_bytes": traffic_limit_bytes, + "period_start_at": created_at, + "hwid_device_limit": _to_int(row.get("device_limit")), + } + metadata = { + "source": SOURCE, + "source_subscription_id": source_id, + "traffic_limit_strategy": str(row.get("traffic_limit_strategy") or ""), + "tag": row.get("tag"), + "internal_squads": [str(item) for item in _listish(row.get("internal_squads"))], + "external_squad": str(row.get("external_squad") or "") or None, + "url": row.get("url"), + "plan_snapshot": plan_snapshot, + } + + if existing: + if self.on_conflict == "skip": + self.summary["subscriptions"]["skipped"] += 1 + else: + for key, value in payload.items(): + self._assign_if_allowed(existing, key, value) + self.summary["subscriptions"]["updated"] += 1 + target_subscription_id = existing.subscription_id + else: + subscription = Subscription(**payload) + self.target.add(subscription) + await self.target.flush() + target_subscription_id = subscription.subscription_id + self.summary["subscriptions"]["created"] += 1 + + await self._upsert_mapping( + entity_type="subscription", + source_id=source_id, + target_table="subscriptions", + target_id=target_subscription_id, + metadata=metadata, + ) + + await self.target.flush() + + async def import_payments(self) -> None: + rows = await self._fetch_rows("transactions", order_by="id") + for row in rows: + user = await self._target_user_for_telegram(row.get("user_telegram_id")) + if not user: + self.summary["payments"]["skipped"] += 1 + continue + + provider_payment_id = f"{SOURCE}:{row.get('payment_id') or row.get('id')}" + existing = ( + await self.target.execute( + select(Payment).where(Payment.provider_payment_id == provider_payment_id) + ) + ).scalar_one_or_none() + + provider = _provider_value(row.get("gateway_type")) + plan_snapshot = _jsonish(row.get("plan_snapshot")) + created_at = _as_utc(row.get("created_at")) + payload = { + "user_id": int(user.user_id), + "provider_payment_id": provider_payment_id, + "provider": provider, + "amount": remnashop_pricing_amount(row.get("pricing")), + "currency": remnashop_pricing_currency(row.get("pricing"), row.get("currency")), + "status": remnashop_transaction_status(row.get("status"), provider), + "description": self._payment_description(row), + "subscription_duration_months": remnashop_months_from_plan_snapshot( + plan_snapshot, + created_at=row.get("created_at"), + expire_at=None, + ), + "sale_mode": remnashop_sale_mode(row.get("purchase_type")), + "tariff_key": remnashop_tariff_key(plan_snapshot, self.tariff_map), + "created_at": created_at, + } + payload = {key: value for key, value in payload.items() if value is not None} + + if existing: + if self.on_conflict == "skip": + self.summary["payments"]["skipped"] += 1 + else: + for key, value in payload.items(): + self._assign_if_allowed(existing, key, value) + self.summary["payments"]["updated"] += 1 + target_payment_id = existing.payment_id + else: + payment = Payment(**payload) + self.target.add(payment) + await self.target.flush() + target_payment_id = payment.payment_id + self.summary["payments"]["created"] += 1 + + await self._upsert_mapping( + entity_type="payment", + source_id=row.get("payment_id") or row.get("id"), + target_table="payments", + target_id=target_payment_id, + metadata={ + "source_transaction_id": row.get("id"), + "is_test": row.get("is_test"), + "purchase_type": str(row.get("purchase_type") or ""), + "gateway_type": str(row.get("gateway_type") or ""), + "plan_snapshot": plan_snapshot, + }, + ) + + await self.target.flush() + + def _payment_description(self, row: dict[str, Any]) -> str: + snapshot = _jsonish(row.get("plan_snapshot")) + plan_name = str(snapshot.get("name") or snapshot.get("tag") or "").strip() + purchase_type = str(row.get("purchase_type") or "").strip().upper() + if plan_name: + return f"Remnashop import: {purchase_type} {plan_name}".strip() + return f"Remnashop import: {purchase_type}".strip() + + async def import_promocodes(self) -> None: + if "promocodes" not in self.tables: + self.summary["promocodes"]["missing_source_table"] += 1 + return + + activation_rows_by_code = await self._source_promocode_activation_rows() + rows = await self._fetch_rows("promocodes", order_by="id") + for row in rows: + code = str(row.get("code") or "").strip() + if not code: + self.summary["promocodes"]["skipped"] += 1 + continue + + bonus_days = self._promo_bonus_days(row) + if bonus_days is None or bonus_days <= 0: + self.summary["promocodes"]["unsupported_reward"] += 1 + continue + + existing = ( + await self.target.execute(select(PromoCode).where(PromoCode.code == code)) + ).scalar_one_or_none() + activations = activation_rows_by_code.get(code, []) + valid_until = None + lifetime_days = _to_int(row.get("lifetime")) + if lifetime_days and _as_utc(row.get("created_at")): + valid_until = _as_utc(row.get("created_at")) + if valid_until: + valid_until = valid_until + timedelta(days=lifetime_days) + + payload = { + "code": code, + "bonus_days": int(bonus_days), + "max_activations": _to_int(row.get("max_activations")) or 1_000_000, + "current_activations": len(activations), + "is_active": bool(row.get("is_active")), + "created_by_admin_id": self.created_by_admin_id, + "created_at": _as_utc(row.get("created_at")), + "valid_until": valid_until, + } + payload = {key: value for key, value in payload.items() if value is not None} + + if existing: + if self.on_conflict == "skip": + self.summary["promocodes"]["skipped"] += 1 + else: + for key, value in payload.items(): + self._assign_if_allowed(existing, key, value) + self.summary["promocodes"]["updated"] += 1 + promo = existing + else: + promo = PromoCode(**payload) + self.target.add(promo) + await self.target.flush() + self.summary["promocodes"]["created"] += 1 + + await self._upsert_mapping( + entity_type="promocode", + source_id=row.get("id") or code, + target_table="promo_codes", + target_id=promo.promo_code_id, + metadata={ + "reward_type": str(row.get("reward_type") or ""), + "reward": row.get("reward"), + "plan": _jsonish(row.get("plan")), + "lifetime": row.get("lifetime"), + }, + ) + await self._import_promocode_activations(promo, activations) + + await self.target.flush() + + async def _source_promocode_activation_rows(self) -> dict[str, list[dict[str, Any]]]: + if "promocode_activations" not in self.tables: + return {} + result = await self.source.execute( + text( + f""" + SELECT a.*, p.code + FROM {_qtable(self.source_schema, "promocode_activations")} a + JOIN {_qtable(self.source_schema, "promocodes")} p + ON p.id = a.promocode_id + ORDER BY a.id + """ + ) + ) + by_code: dict[str, list[dict[str, Any]]] = defaultdict(list) + for row in result.mappings().all(): + mapping = _as_mapping(row) + code = str(mapping.get("code") or "").strip() + if code: + by_code[code].append(mapping) + return by_code + + def _promo_bonus_days(self, row: dict[str, Any]) -> Optional[int]: + reward_type = str(row.get("reward_type") or "").strip().upper() + if reward_type == "DURATION": + return _to_int(row.get("reward")) + if reward_type == "SUBSCRIPTION": + plan = _jsonish(row.get("plan")) + return ( + _to_int(plan.get("duration_days")) + or _to_int(plan.get("days")) + or _to_int(row.get("reward")) + ) + return None + + async def _import_promocode_activations( + self, + promo: PromoCode, + activations: Iterable[dict[str, Any]], + ) -> None: + for activation in activations: + user = await self._target_user_for_telegram(activation.get("user_telegram_id")) + if not user: + self.summary["promocodes"]["activation_skipped"] += 1 + continue + stmt = ( + pg_insert(PromoCodeActivation) + .values( + promo_code_id=promo.promo_code_id, + user_id=user.user_id, + activated_at=_as_utc(activation.get("activated_at")) + or datetime.now(timezone.utc), + ) + .on_conflict_do_nothing( + index_elements=[ + PromoCodeActivation.promo_code_id, + PromoCodeActivation.user_id, + ] + ) + ) + await self.target.execute(stmt) + self.summary["promocodes"]["activation_imported"] += 1 + + async def import_settings(self) -> None: + source_settings = await self._fetch_one("settings") + plans = ( + await self._fetch_rows("plans", order_by="order_index") + if "plans" in self.tables + else [] + ) + notes = { + "default_currency": ( + source_settings.get("default_currency") if source_settings else None + ), + "settings": { + key: source_settings.get(key) + for key in ("access", "requirements", "notifications", "referral", "menu") + if source_settings and source_settings.get(key) is not None + }, + "plans_count": len(plans), + "plans": [ + { + "id": plan.get("id"), + "name": plan.get("name"), + "type": str(plan.get("type") or ""), + "traffic_limit": plan.get("traffic_limit"), + "device_limit": plan.get("device_limit"), + "tag": plan.get("tag"), + } + for plan in plans[:100] + ], + } + await self._upsert_mapping( + entity_type="settings", + source_id="singleton", + target_table="app_setting_overrides", + target_id="MIGRATION_REMNASHOP_NOTES", + metadata=notes, + ) + self.summary["settings"]["captured"] += 1 + + async def _write_admin_overrides(self) -> None: + now = datetime.now(timezone.utc).isoformat() + plain_summary = self._plain_summary() + await self._upsert_setting_override( + "MIGRATION_REMNASHOP_REFERRAL_CODE_COMPAT_ENABLED", + True, + ) + await self._upsert_setting_override( + "MIGRATION_REMNASHOP_PROMO_CODE_COMPAT_ENABLED", + "promocodes" in self.tables, + ) + await self._upsert_setting_override("MIGRATION_REMNASHOP_IMPORTED_AT", now) + await self._upsert_setting_override( + "MIGRATION_REMNASHOP_NOTES", + _json_dumps(plain_summary), + ) + self.summary["settings"]["admin_overrides_written"] += 1 + + +def parse_only(value: str) -> set[str]: + if not value: + return set() + return {item.strip().lower() for item in value.split(",") if item.strip()} + + +def parse_tariff_map(value: Optional[str]) -> dict[str, str]: + if not value: + return {} + path = Path(value) + raw = path.read_text(encoding="utf-8") if path.exists() else value + decoded = json.loads(raw) + if not isinstance(decoded, dict): + raise ValueError("--tariff-map-json must be a JSON object or a path to one") + return {str(key): str(mapped) for key, mapped in decoded.items()} + + +def build_arg_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description="Import legacy bot data into this shop.") + parser.add_argument("--source-type", choices=[SOURCE], default=SOURCE) + parser.add_argument("--source-dsn", required=True) + parser.add_argument("--source-schema", default="public") + parser.add_argument("--target-dsn") + parser.add_argument( + "--only", + default="all", + help=( + "Comma-separated sections: " + "all,users,referrals,subscriptions,payments,promocodes,settings" + ), + ) + parser.add_argument( + "--on-conflict", + choices=["merge", "skip", "overwrite"], + default="merge", + ) + parser.add_argument("--dry-run", action="store_true") + parser.add_argument("--created-by-admin-id", type=int, default=0) + parser.add_argument( + "--tariff-map-json", + help="JSON object or path mapping remnashop plan id/name/tag to local tariff_key.", + ) + parser.add_argument( + "--no-admin-compat-overrides", + action="store_true", + help="Do not enable migration compatibility toggles in admin settings.", + ) + return parser + + +async def _prepare_target_schema(engine: Any) -> None: + async with engine.begin() as connection: + await connection.run_sync(Base.metadata.create_all) + await connection.run_sync(run_database_migrations) + + +async def run_import(args: argparse.Namespace) -> dict[str, Any]: + settings = Settings() + source_engine = create_async_engine(normalize_async_postgres_dsn(args.source_dsn)) + target_engine = create_async_engine( + normalize_async_postgres_dsn(args.target_dsn or settings.DATABASE_URL) + ) + await _prepare_target_schema(target_engine) + + session_factory = async_sessionmaker( + bind=target_engine, + class_=AsyncSession, + expire_on_commit=False, + autocommit=False, + autoflush=False, + ) + + async with source_engine.connect() as source, session_factory() as target: + importer = RemnashopImporter( + source=source, + target=target, + source_schema=args.source_schema, + only=parse_only(args.only), + on_conflict=args.on_conflict, + dry_run=bool(args.dry_run), + created_by_admin_id=args.created_by_admin_id, + tariff_map=parse_tariff_map(args.tariff_map_json), + write_admin_compat_overrides=not args.no_admin_compat_overrides, + ) + summary = await importer.run() + if args.dry_run: + await target.rollback() + else: + await target.commit() + + await source_engine.dispose() + await target_engine.dispose() + return summary + + +def main() -> None: + logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s") + args = build_arg_parser().parse_args() + summary = asyncio.run(run_import(args)) + print(_json_dumps(summary)) + + +if __name__ == "__main__": + main() diff --git a/frontend/src/admin/sections/SettingsSection.svelte b/frontend/src/admin/sections/SettingsSection.svelte index 3443016..8d151fb 100644 --- a/frontend/src/admin/sections/SettingsSection.svelte +++ b/frontend/src/admin/sections/SettingsSection.svelte @@ -348,6 +348,7 @@ devices: "Устройства", subscription_guides: "Connection guides", system: "Система", + migrations: "Миграции", }; return adminText(`settings_section_${id}`, {}, map[id] || id); } diff --git a/locales/en.json b/locales/en.json index cf3ffff..c6ae75d 100644 --- a/locales/en.json +++ b/locales/en.json @@ -1108,11 +1108,13 @@ "admin_settings_section_devices": "Devices", "admin_settings_section_support": "Support", "admin_settings_section_system": "System", + "admin_settings_section_migrations": "Migrations", "admin_settings_field_telemetry_enabled_label": "Anonymous install analytics", "admin_settings_field_telemetry_enabled_description": "Sends one anonymous heartbeat per day (version, OS, locale, user-count range). No personal data, tokens or domains. Helps gauge how many installs are active and which versions are in use. Toggling this off takes effect without a restart.", "admin_settings_subsection_common": "Common", "admin_settings_subsection_checkout": "Checkout", "admin_settings_subsection_remnawave": "Remnawave", + "admin_settings_subsection_remnashop": "Remnashop", "admin_settings_subsection_telegram_stars": "Telegram Stars", "admin_settings_subsection_yookassa": "YooKassa", "admin_settings_subsection_freekassa": "FreeKassa", @@ -1127,6 +1129,14 @@ "admin_settings_provider_webhook_base_missing": "Set WEBHOOK_BASE_URL in .env to show the full URL for {path}.", "admin_settings_provider_admin_only_label": "Only for admins", "admin_settings_provider_admin_only_description": "Shows this provider only to admins. Webhooks and payment status handling remain active for test payments.", + "admin_settings_field_migration_remnashop_referral_code_compat_enabled_label": "Remnashop ref-code compatibility", + "admin_settings_field_migration_remnashop_referral_code_compat_enabled_description": "Allows old Remnashop referral codes to resolve without changing their case or format. The importer enables it automatically.", + "admin_settings_field_migration_remnashop_promo_code_compat_enabled_label": "Remnashop promo-code compatibility", + "admin_settings_field_migration_remnashop_promo_code_compat_enabled_description": "Looks up promo codes in the original case first, then falls back to the current rules. Regular promo codes keep working as before.", + "admin_settings_field_migration_remnashop_imported_at_label": "Remnashop import date", + "admin_settings_field_migration_remnashop_imported_at_description": "Operational marker for the latest Remnashop import, stored as an ISO timestamp.", + "admin_settings_field_migration_remnashop_notes_label": "Remnashop import notes", + "admin_settings_field_migration_remnashop_notes_description": "Short importer summary: migrated entities and enabled compatibility modes.", "admin_settings_validation_errors": "Errors: {errors}", "admin_settings_save_error": "Error: {error}", "admin_sync_started": "Synchronization started", diff --git a/locales/ru.json b/locales/ru.json index 4ce0660..e7d5047 100644 --- a/locales/ru.json +++ b/locales/ru.json @@ -1108,11 +1108,13 @@ "admin_settings_section_devices": "Устройства", "admin_settings_section_support": "Поддержка", "admin_settings_section_system": "Система", + "admin_settings_section_migrations": "Миграции", "admin_settings_field_telemetry_enabled_label": "Анонимная статистика установки", "admin_settings_field_telemetry_enabled_description": "Раз в сутки отправляет обезличенный сигнал: версия, ОС, локаль и число пользователей в виде диапазона. Без персональных данных, токенов и доменов. Помогает оценить число активных установок и используемые версии. Отключение применяется без перезапуска.", "admin_settings_subsection_common": "Общие", "admin_settings_subsection_checkout": "Оформление оплаты", "admin_settings_subsection_remnawave": "Remnawave", + "admin_settings_subsection_remnashop": "Remnashop", "admin_settings_subsection_telegram_stars": "Telegram Stars", "admin_settings_subsection_yookassa": "YooKassa", "admin_settings_subsection_freekassa": "FreeKassa", @@ -1127,6 +1129,14 @@ "admin_settings_provider_webhook_base_missing": "Укажите WEBHOOK_BASE_URL в .env, чтобы увидеть полный адрес для {path}.", "admin_settings_provider_admin_only_label": "Только для админов", "admin_settings_provider_admin_only_description": "Показывает провайдер только администраторам. Вебхуки и обработка статусов остаются активными для тестовых платежей.", + "admin_settings_field_migration_remnashop_referral_code_compat_enabled_label": "Совместимость ref-кодов Remnashop", + "admin_settings_field_migration_remnashop_referral_code_compat_enabled_description": "Разрешает вход по старым ref-кодам Remnashop без изменения их регистра и формата. Включается импортёром автоматически.", + "admin_settings_field_migration_remnashop_promo_code_compat_enabled_label": "Совместимость промокодов Remnashop", + "admin_settings_field_migration_remnashop_promo_code_compat_enabled_description": "Ищет промокоды сначала в исходном регистре, затем по текущим правилам. Обычные промокоды продолжают работать как раньше.", + "admin_settings_field_migration_remnashop_imported_at_label": "Дата импорта Remnashop", + "admin_settings_field_migration_remnashop_imported_at_description": "Служебная отметка последнего импорта Remnashop в ISO-формате.", + "admin_settings_field_migration_remnashop_notes_label": "Заметки импорта Remnashop", + "admin_settings_field_migration_remnashop_notes_description": "Краткая сводка импортёра: какие сущности перенесены и какие совместимые режимы включены.", "admin_settings_validation_errors": "Ошибки: {errors}", "admin_settings_save_error": "Ошибка: {error}", "admin_sync_started": "Синхронизация запущена", diff --git a/tests/test_admin_settings_manifest_i18n.py b/tests/test_admin_settings_manifest_i18n.py index 2144100..21ecbed 100644 --- a/tests/test_admin_settings_manifest_i18n.py +++ b/tests/test_admin_settings_manifest_i18n.py @@ -44,6 +44,13 @@ BACKUP_SETTINGS = ( "BACKUP_COMPOSE_ENABLED", ) +REMNASHOP_MIGRATION_SETTINGS = ( + "MIGRATION_REMNASHOP_REFERRAL_CODE_COMPAT_ENABLED", + "MIGRATION_REMNASHOP_PROMO_CODE_COMPAT_ENABLED", + "MIGRATION_REMNASHOP_IMPORTED_AT", + "MIGRATION_REMNASHOP_NOTES", +) + ADMIN_TARIFF_SETTINGS_PAGE_KEYS = { "admin_tariffs_trial_title", "admin_tariffs_trial_subtitle", @@ -204,6 +211,27 @@ def test_backup_settings_i18n_keys_exist(): assert field["i18n_description_key"] in messages +def test_remnashop_migration_settings_i18n_keys_exist(): + manifest = _manifest_by_key() + + for setting_key in REMNASHOP_MIGRATION_SETTINGS: + field = manifest[setting_key] + assert field["section"] == "migrations" + assert field["section_order"] == 13 + assert field["subsection"] == "Remnashop" + assert field["i18n_subsection_key"] == "admin_settings_subsection_remnashop" + + for language in ("ru", "en"): + messages = _locale(language) + + assert "admin_settings_section_migrations" in messages + assert "admin_settings_subsection_remnashop" in messages + for setting_key in REMNASHOP_MIGRATION_SETTINGS: + field = manifest[setting_key] + assert field["i18n_label_key"] in messages + assert field["i18n_description_key"] in messages + + def test_backup_required_numeric_settings_reject_empty_values(): with pytest.raises(ValueError): coerce_value(get_field_by_key("BACKUP_INTERVAL_SECONDS"), "") diff --git a/tests/test_remnashop_import.py b/tests/test_remnashop_import.py new file mode 100644 index 0000000..fb4a16b --- /dev/null +++ b/tests/test_remnashop_import.py @@ -0,0 +1,44 @@ +from datetime import datetime, timezone + +from scripts.import_legacy import ( + remnashop_months_from_plan_snapshot, + remnashop_pricing_amount, + remnashop_pricing_currency, + remnashop_sale_mode, + remnashop_traffic_gb_to_bytes, + remnashop_transaction_status, +) + + +def test_remnashop_pricing_helpers_read_final_amount_and_currency(): + pricing = {"final_amount": "199.50", "currency": "rub"} + + assert remnashop_pricing_amount(pricing) == 199.5 + assert remnashop_pricing_currency(pricing) == "RUB" + + +def test_remnashop_traffic_limit_is_converted_from_gib(): + assert remnashop_traffic_gb_to_bytes(10) == 10 * 1024**3 + assert remnashop_traffic_gb_to_bytes(None) is None + + +def test_remnashop_status_and_sale_mode_mapping_matches_current_payment_model(): + assert remnashop_transaction_status("COMPLETED", "YOOKASSA") == "succeeded" + assert remnashop_transaction_status("PENDING", "WATA") == "pending_wata" + assert remnashop_transaction_status("CANCELED", "WATA") == "canceled" + assert remnashop_sale_mode("NEW") == "subscription" + assert remnashop_sale_mode("RENEW") == "subscription" + assert remnashop_sale_mode("CHANGE") == "tariff_upgrade" + + +def test_remnashop_plan_months_prefers_snapshot_then_dates(): + assert remnashop_months_from_plan_snapshot({"duration_days": 90}) == 3 + assert remnashop_months_from_plan_snapshot({"months": 12}) == 12 + assert ( + remnashop_months_from_plan_snapshot( + {}, + created_at=datetime(2026, 1, 1, tzinfo=timezone.utc), + expire_at=datetime(2026, 4, 1, tzinfo=timezone.utc), + ) + == 3 + ) diff --git a/tests/test_support_migration.py b/tests/test_support_migration.py index 5fa5292..2e5fed5 100644 --- a/tests/test_support_migration.py +++ b/tests/test_support_migration.py @@ -1,5 +1,11 @@ from db.migrator import MIGRATIONS -from db.models import SupportTicket, SupportTicketMessage, User +from db.models import ( + LegacyImportMapping, + LegacyReferralCode, + SupportTicket, + SupportTicketMessage, + User, +) def test_support_migration_is_registered_after_existing_revisions(): @@ -39,3 +45,18 @@ def test_trial_eligibility_reset_migration_and_model_are_registered(): "0032_add_telegram_notification_status" ) assert "trial_eligibility_reset_at" in User.__table__.columns + + +def test_legacy_import_compatibility_migration_and_models_are_registered(): + ids = [migration.id for migration in MIGRATIONS] + + assert "0034_add_legacy_import_compatibility" in ids + assert ids.index("0034_add_legacy_import_compatibility") > ids.index( + "0033_add_trial_eligibility_reset_marker" + ) + assert User.__table__.columns["referral_code"].type.length == 64 + assert LegacyReferralCode.__tablename__ == "legacy_referral_codes" + assert LegacyImportMapping.__tablename__ == "legacy_import_mappings" + assert "uq_legacy_referral_source_code" in { + constraint.name for constraint in LegacyReferralCode.__table__.constraints + } diff --git a/tests/test_user_dal.py b/tests/test_user_dal.py index 5b8bd30..8849ba9 100644 --- a/tests/test_user_dal.py +++ b/tests/test_user_dal.py @@ -230,6 +230,8 @@ class UserDalMergeTests(unittest.IsolatedAsyncioTestCase): ) self.assertIn("support_ticket_messages", update_tables) self.assertIn("email_verification_codes", delete_tables) + self.assertIn("legacy_referral_codes", delete_tables) + self.assertIn("legacy_import_mappings", delete_tables) session.delete.assert_awaited_once_with(user) session.flush.assert_awaited_once() @@ -366,6 +368,8 @@ class UserDalMergeTests(unittest.IsolatedAsyncioTestCase): self.assertIn("payments", update_tables) self.assertIn("promo_code_activations", update_tables) self.assertIn("user_payment_methods", update_tables) + self.assertIn("legacy_referral_codes", update_tables) + self.assertIn("legacy_import_mappings", update_tables) self.assertIn("message_logs", update_tables) self.assertIn("users", update_tables) self.assertIn("user_payment_methods", delete_tables)