"""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()