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.
This commit is contained in:
3252a8
2026-06-02 00:13:54 +03:00
parent 56796d9f22
commit 5e47a1496b
18 changed files with 1731 additions and 63 deletions
+35 -13
View File
@@ -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(
+93 -6
View File
@@ -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))
+78
View File
@@ -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,
),
]
+32 -1
View File
@@ -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"