diff --git a/.env.example b/.env.example
index 0e59912..48d0a3b 100644
--- a/.env.example
+++ b/.env.example
@@ -176,6 +176,9 @@ LOG_PROMO_ACTIVATIONS=True #
LOG_TRIAL_ACTIVATIONS=True # Log trial activations
LOG_SUSPICIOUS_ACTIVITY=True # Log suspicious activity
+# Discount promo reservation timeout
+DISCOUNT_PROMO_PAYMENT_TIMEOUT_MINUTES=10 # Minutes to keep discount promo reservation before it expires
+
# Embedded mode thumbnails. Please don't touch this if you don't know what it is.
INLINE_REFERRAL_THUMBNAIL_URL=https://cdn-icons-png.flaticon.com/512/1077/1077114.png
INLINE_USER_STATS_THUMBNAIL_URL=https://cdn-icons-png.flaticon.com/512/681/681494.png
diff --git a/alembic/versions/0002_active_discount_expires_at.py b/alembic/versions/0002_active_discount_expires_at.py
new file mode 100644
index 0000000..4a76a52
--- /dev/null
+++ b/alembic/versions/0002_active_discount_expires_at.py
@@ -0,0 +1,86 @@
+"""add active discount expiration
+
+Revision ID: 0002_active_discount_expires_at
+Revises: 0001_initial_schema
+Create Date: 2026-02-08 00:00:01.000000
+
+"""
+
+from typing import Sequence, Union
+
+from alembic import op, context
+import sqlalchemy as sa
+
+
+# revision identifiers, used by Alembic.
+revision: str = "0002_active_discount_expires_at"
+down_revision: Union[str, Sequence[str], None] = "0001_initial_schema"
+branch_labels: Union[str, Sequence[str], None] = None
+depends_on: Union[str, Sequence[str], None] = None
+
+
+_INDEX_NAME = "idx_active_discounts_expires_at"
+
+
+def upgrade() -> None:
+ if context.is_offline_mode():
+ op.add_column(
+ "active_discounts",
+ sa.Column("expires_at", sa.DateTime(timezone=True), nullable=True),
+ )
+ op.execute(
+ sa.text(
+ "UPDATE active_discounts "
+ "SET expires_at = COALESCE(activated_at, NOW()) + INTERVAL '10 minutes' "
+ "WHERE expires_at IS NULL"
+ )
+ )
+ op.alter_column("active_discounts", "expires_at", nullable=False)
+ op.create_index(_INDEX_NAME, "active_discounts", ["expires_at"], unique=False)
+ return
+
+ bind = op.get_bind()
+ inspector = sa.inspect(bind)
+
+ if not inspector.has_table("active_discounts"):
+ return
+
+ columns = {column["name"] for column in inspector.get_columns("active_discounts")}
+ if "expires_at" not in columns:
+ op.add_column(
+ "active_discounts",
+ sa.Column("expires_at", sa.DateTime(timezone=True), nullable=True),
+ )
+ op.execute(
+ sa.text(
+ "UPDATE active_discounts "
+ "SET expires_at = COALESCE(activated_at, NOW()) + INTERVAL '10 minutes' "
+ "WHERE expires_at IS NULL"
+ )
+ )
+ op.alter_column("active_discounts", "expires_at", nullable=False)
+
+ indexes = {index["name"] for index in inspector.get_indexes("active_discounts")}
+ if _INDEX_NAME not in indexes:
+ op.create_index(_INDEX_NAME, "active_discounts", ["expires_at"], unique=False)
+
+
+def downgrade() -> None:
+ if context.is_offline_mode():
+ op.drop_index(_INDEX_NAME, table_name="active_discounts")
+ op.drop_column("active_discounts", "expires_at")
+ return
+
+ bind = op.get_bind()
+ inspector = sa.inspect(bind)
+
+ if not inspector.has_table("active_discounts"):
+ return
+
+ indexes = {index["name"] for index in inspector.get_indexes("active_discounts")}
+ if _INDEX_NAME in indexes:
+ op.drop_index(_INDEX_NAME, table_name="active_discounts")
+
+ columns = {column["name"] for column in inspector.get_columns("active_discounts")}
+ if "expires_at" in columns:
+ op.drop_column("active_discounts", "expires_at")
diff --git a/bot/main_bot.py b/bot/main_bot.py
index ae15640..93dcdda 100644
--- a/bot/main_bot.py
+++ b/bot/main_bot.py
@@ -144,6 +144,18 @@ async def on_startup_configured(dispatcher: Dispatcher):
except Exception as e:
logging.error(f"STARTUP: Failed to initialize message queue manager: {e}", exc_info=True)
+ # Initialize promo discount expiration worker
+ try:
+ promo_code_service: Optional[PromoCodeService] = dispatcher.get("promo_code_service")
+ if promo_code_service:
+ await promo_code_service.setup_discount_expiration_worker(async_session_factory)
+ logging.info("STARTUP: Promo discount expiration worker initialized")
+ except Exception as e:
+ logging.error(
+ f"STARTUP: Failed to initialize promo discount expiration worker: {e}",
+ exc_info=True,
+ )
+
# Automatic sync on startup
try:
logging.info("STARTUP: Running automatic panel sync...")
diff --git a/bot/services/promo_code_service.py b/bot/services/promo_code_service.py
index 3fb90ad..0d8475f 100644
--- a/bot/services/promo_code_service.py
+++ b/bot/services/promo_code_service.py
@@ -1,13 +1,14 @@
import logging
-from datetime import datetime, timezone
+import asyncio
+from datetime import datetime, timezone, timedelta
from sqlalchemy.ext.asyncio import AsyncSession
-from typing import Optional, Tuple, Dict
+from typing import Optional, Tuple
from aiogram import Bot
+from sqlalchemy.orm import sessionmaker
from config.settings import Settings
from db.dal import promo_code_dal, user_dal, active_discount_dal, payment_dal
-from db.models import PromoCode, User
from .subscription_service import SubscriptionService
from bot.middlewares.i18n import JsonI18n
@@ -23,6 +24,117 @@ class PromoCodeService:
self.subscription_service = subscription_service
self.bot = bot
self.i18n = i18n
+ self.discount_payment_timeout_minutes = max(
+ 1,
+ int(getattr(settings, "DISCOUNT_PROMO_PAYMENT_TIMEOUT_MINUTES", 10) or 10),
+ )
+ self._discount_expiration_task: Optional[asyncio.Task] = None
+ self._async_session_factory: Optional[sessionmaker] = None
+
+ async def setup_discount_expiration_worker(
+ self,
+ async_session_factory: sessionmaker,
+ ) -> None:
+ """Attach DB session factory and start background cleanup loop."""
+ self._async_session_factory = async_session_factory
+ if self._discount_expiration_task and not self._discount_expiration_task.done():
+ return
+ self._discount_expiration_task = asyncio.create_task(
+ self._discount_expiration_loop(),
+ name="PromoDiscountExpirationLoop",
+ )
+ logging.info("PromoCodeService: started discount expiration background worker.")
+
+ async def close(self) -> None:
+ """Gracefully stop background workers."""
+ if not self._discount_expiration_task:
+ return
+ self._discount_expiration_task.cancel()
+ try:
+ await self._discount_expiration_task
+ except asyncio.CancelledError:
+ pass
+ except Exception:
+ logging.exception("PromoCodeService: failed while stopping expiration worker")
+ finally:
+ self._discount_expiration_task = None
+
+ async def _discount_expiration_loop(self) -> None:
+ """Periodically clears expired discount reservations and notifies users."""
+ while True:
+ try:
+ if not self._async_session_factory:
+ await asyncio.sleep(30)
+ continue
+
+ await self._process_expired_discounts_once()
+ except asyncio.CancelledError:
+ logging.info("PromoCodeService: discount expiration loop cancelled.")
+ raise
+ except Exception:
+ logging.exception("PromoCodeService: unhandled error in discount expiration loop")
+
+ await asyncio.sleep(30)
+
+ async def _process_expired_discounts_once(self) -> None:
+ if not self._async_session_factory:
+ return
+
+ now_utc = datetime.now(timezone.utc)
+ async with self._async_session_factory() as session:
+ expired_discounts = await active_discount_dal.get_expired_active_discounts(
+ session,
+ now=now_utc,
+ limit=100,
+ )
+ if not expired_discounts:
+ return
+
+ for expired in expired_discounts:
+ cleared = await active_discount_dal.clear_active_discount_if_matches(
+ session,
+ user_id=expired.user_id,
+ promo_code_id=expired.promo_code_id,
+ expires_at_lte=now_utc,
+ )
+ if not cleared:
+ continue
+
+ await promo_code_dal.decrement_promo_code_usage(session, expired.promo_code_id)
+
+ db_user = await user_dal.get_user_by_id(session, expired.user_id)
+ user_lang = (
+ db_user.language_code
+ if db_user and db_user.language_code
+ else self.settings.DEFAULT_LANGUAGE
+ )
+ promo = await promo_code_dal.get_promo_code_by_id(session, expired.promo_code_id)
+ promo_code = promo.code if promo else ""
+ message_text = self.i18n.gettext(
+ user_lang,
+ "discount_promo_expired_need_reactivate",
+ code_part=(f" ({promo_code})" if promo_code else ""),
+ )
+
+ try:
+ await self.bot.send_message(
+ chat_id=expired.user_id,
+ text=message_text,
+ parse_mode="HTML",
+ )
+ except Exception:
+ logging.exception(
+ "Failed to send discount expiration message to user %s",
+ expired.user_id,
+ )
+
+ logging.info(
+ "Expired discount reservation removed: user=%s, promo=%s",
+ expired.user_id,
+ expired.promo_code_id,
+ )
+
+ await session.commit()
async def apply_promo_code(
self,
@@ -99,7 +211,26 @@ class PromoCodeService:
code_input_upper = code_input.strip().upper()
# Check if user already has an active discount
- existing_discount = await active_discount_dal.get_active_discount(session, user_id)
+ existing_discount = await active_discount_dal.get_active_discount(
+ session,
+ user_id,
+ include_expired=True,
+ )
+ if existing_discount:
+ now_utc = datetime.now(timezone.utc)
+ if existing_discount.expires_at <= now_utc:
+ cleared = await active_discount_dal.clear_active_discount_if_expired(
+ session,
+ user_id,
+ now=now_utc,
+ )
+ if cleared:
+ await promo_code_dal.decrement_promo_code_usage(
+ session,
+ existing_discount.promo_code_id,
+ )
+ existing_discount = None
+
if existing_discount:
# Get the promo code for the existing discount
existing_promo = await promo_code_dal.get_promo_code_by_id(
@@ -128,21 +259,37 @@ class PromoCodeService:
if existing_activation:
return False, _("promo_code_already_used_by_user", code=code_input_upper)
- # Set active discount
+ # Reserve discount for limited time and count activation immediately
+ expires_at = datetime.now(timezone.utc) + timedelta(
+ minutes=self.discount_payment_timeout_minutes,
+ )
active_discount = await active_discount_dal.set_active_discount(
session,
user_id=user_id,
promo_code_id=promo_data.promo_code_id,
- discount_percentage=promo_data.discount_percentage
+ discount_percentage=promo_data.discount_percentage,
+ expires_at=expires_at,
)
if not active_discount:
# This shouldn't happen since we checked above, but just in case
return False, _("error_applying_promo_discount")
+ promo_incremented = await promo_code_dal.increment_promo_code_usage(
+ session,
+ promo_data.promo_code_id,
+ )
+ if not promo_incremented:
+ await active_discount_dal.clear_active_discount_if_matches(
+ session,
+ user_id=user_id,
+ promo_code_id=promo_data.promo_code_id,
+ )
+ return False, _("promo_code_not_found_or_not_discount", code=code_input_upper)
+
logging.info(
f"Discount promo code {code_input_upper} activated for user {user_id}: "
- f"{promo_data.discount_percentage}% off"
+ f"{promo_data.discount_percentage}% off until {expires_at.isoformat()}"
)
return True, promo_data.discount_percentage
@@ -155,10 +302,28 @@ class PromoCodeService:
Get user's active discount if any.
Returns: (discount_percentage, promo_code) or None
"""
- active_discount = await active_discount_dal.get_active_discount(session, user_id)
+ active_discount = await active_discount_dal.get_active_discount(
+ session,
+ user_id,
+ include_expired=True,
+ )
if not active_discount:
return None
+ now_utc = datetime.now(timezone.utc)
+ if active_discount.expires_at <= now_utc:
+ cleared = await active_discount_dal.clear_active_discount_if_expired(
+ session,
+ user_id,
+ now=now_utc,
+ )
+ if cleared:
+ await promo_code_dal.decrement_promo_code_usage(
+ session,
+ active_discount.promo_code_id,
+ )
+ return None
+
# Fetch promo code for code string
promo = await promo_code_dal.get_promo_code_by_id(
session, active_discount.promo_code_id
@@ -175,7 +340,9 @@ class PromoCodeService:
f"Promo code {promo.code} expired (valid_until: {promo.valid_until}). "
f"Clearing active discount for user {user_id}"
)
- await active_discount_dal.clear_active_discount(session, user_id)
+ cleared = await active_discount_dal.clear_active_discount(session, user_id)
+ if cleared:
+ await promo_code_dal.decrement_promo_code_usage(session, promo.promo_code_id)
return None
return (active_discount.discount_percentage, promo.code)
@@ -206,7 +373,7 @@ class PromoCodeService:
payment_id: int
) -> bool:
"""
- Consume active discount: link activation to payment, increment usage, clear active discount.
+ Consume active discount: link reservation to payment and clear active discount.
Call this AFTER successful payment.
"""
payment_record = await payment_dal.get_payment_by_db_id(session, payment_id)
@@ -230,7 +397,11 @@ class PromoCodeService:
)
return False
- active_discount = await active_discount_dal.get_active_discount(session, user_id)
+ active_discount = await active_discount_dal.get_active_discount(
+ session,
+ user_id,
+ include_expired=True,
+ )
if active_discount and active_discount.promo_code_id != promo_code_id:
logging.info(
"Active discount promo %s differs from payment promo %s; leaving active discount intact.",
@@ -239,6 +410,27 @@ class PromoCodeService:
)
active_discount = None
+ now_utc = datetime.now(timezone.utc)
+ if (
+ active_discount
+ and active_discount.promo_code_id == promo_code_id
+ and active_discount.expires_at <= now_utc
+ ):
+ logging.info(
+ "Discount reservation expired before payment consumption (user=%s, promo=%s)",
+ user_id,
+ promo_code_id,
+ )
+ cleared = await active_discount_dal.clear_active_discount_if_matches(
+ session,
+ user_id=user_id,
+ promo_code_id=promo_code_id,
+ expires_at_lte=now_utc,
+ )
+ if cleared:
+ await promo_code_dal.decrement_promo_code_usage(session, promo_code_id)
+ return False
+
existing_activation = await promo_code_dal.get_user_activation_for_promo(
session, promo_code_id, user_id
)
@@ -269,17 +461,6 @@ class PromoCodeService:
)
return False
- promo_incremented = await promo_code_dal.increment_promo_code_usage(
- session, promo_code_id, allow_overflow=True
- )
- if not promo_incremented:
- logging.error(
- "Failed to increment discount usage for user %s, promo %s.",
- user_id,
- promo_code_id,
- )
- return False
-
if active_discount and active_discount.promo_code_id == promo_code_id:
await active_discount_dal.clear_active_discount(session, user_id)
diff --git a/config/settings.py b/config/settings.py
index f6a0be2..068bced 100644
--- a/config/settings.py
+++ b/config/settings.py
@@ -652,6 +652,10 @@ class Settings(BaseSettings):
LOG_PROMO_ACTIVATIONS: bool = Field(default=True, description="Send notifications for promo code activations")
LOG_TRIAL_ACTIVATIONS: bool = Field(default=True, description="Send notifications for trial activations")
LOG_SUSPICIOUS_ACTIVITY: bool = Field(default=True, description="Send notifications for suspicious promo attempts")
+ DISCOUNT_PROMO_PAYMENT_TIMEOUT_MINUTES: int = Field(
+ default=10,
+ description="How long a discount promo reservation is kept before user payment",
+ )
model_config = SettingsConfigDict(env_file='.env',
env_file_encoding='utf-8',
diff --git a/db/dal/active_discount_dal.py b/db/dal/active_discount_dal.py
index 5b99c51..e9af825 100644
--- a/db/dal/active_discount_dal.py
+++ b/db/dal/active_discount_dal.py
@@ -1,38 +1,44 @@
import logging
-from typing import Optional
+from typing import Optional, List
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.future import select
from sqlalchemy import delete
from datetime import datetime, timezone
-from db.models import ActiveDiscount, PromoCode
+from db.models import ActiveDiscount
async def set_active_discount(
session: AsyncSession,
user_id: int,
promo_code_id: int,
- discount_percentage: int
+ discount_percentage: int,
+ expires_at: datetime,
) -> Optional[ActiveDiscount]:
"""
Set active discount for user.
Returns None if user already has an active discount (enforce one-at-a-time rule).
"""
- # Check if user already has an active discount
- existing = await get_active_discount(session, user_id)
- if existing:
+ now_utc = datetime.now(timezone.utc)
+
+ existing = await get_active_discount(session, user_id, include_expired=True)
+ if existing and existing.expires_at > now_utc:
logging.warning(
f"User {user_id} already has active discount (promo_code_id: {existing.promo_code_id}). "
f"Cannot activate new discount {promo_code_id}."
)
return None
+ if existing and existing.expires_at <= now_utc:
+ await clear_active_discount_if_expired(session, user_id, now=now_utc)
+
# Create new active discount
new_discount = ActiveDiscount(
user_id=user_id,
promo_code_id=promo_code_id,
discount_percentage=discount_percentage,
- activated_at=datetime.now(timezone.utc)
+ activated_at=now_utc,
+ expires_at=expires_at,
)
session.add(new_discount)
await session.flush()
@@ -46,10 +52,14 @@ async def set_active_discount(
async def get_active_discount(
session: AsyncSession,
- user_id: int
+ user_id: int,
+ include_expired: bool = False,
) -> Optional[ActiveDiscount]:
"""Get active discount for user if exists."""
+ now_utc = datetime.now(timezone.utc)
stmt = select(ActiveDiscount).where(ActiveDiscount.user_id == user_id)
+ if not include_expired:
+ stmt = stmt.where(ActiveDiscount.expires_at > now_utc)
result = await session.execute(stmt)
return result.scalar_one_or_none()
@@ -71,6 +81,71 @@ async def clear_active_discount(
return cleared
+async def clear_active_discount_if_expired(
+ session: AsyncSession,
+ user_id: int,
+ now: Optional[datetime] = None,
+) -> bool:
+ """
+ Clear active discount for user only when it has already expired.
+ """
+ now_utc = now or datetime.now(timezone.utc)
+ stmt = delete(ActiveDiscount).where(
+ ActiveDiscount.user_id == user_id,
+ ActiveDiscount.expires_at <= now_utc,
+ )
+ result = await session.execute(stmt)
+ await session.flush()
+ cleared = result.rowcount > 0
+ if cleared:
+ logging.info("Expired active discount cleared for user %s", user_id)
+ return cleared
+
+
+async def clear_active_discount_if_matches(
+ session: AsyncSession,
+ user_id: int,
+ promo_code_id: Optional[int] = None,
+ expires_at_lte: Optional[datetime] = None,
+) -> bool:
+ """
+ Clear active discount for user only when additional constraints match.
+ """
+ conditions = [ActiveDiscount.user_id == user_id]
+ if promo_code_id is not None:
+ conditions.append(ActiveDiscount.promo_code_id == promo_code_id)
+ if expires_at_lte is not None:
+ conditions.append(ActiveDiscount.expires_at <= expires_at_lte)
+
+ stmt = delete(ActiveDiscount).where(*conditions)
+ result = await session.execute(stmt)
+ await session.flush()
+ cleared = result.rowcount > 0
+ if cleared:
+ logging.info(
+ "Active discount cleared for user %s by constrained cleanup.",
+ user_id,
+ )
+ return cleared
+
+
+async def get_expired_active_discounts(
+ session: AsyncSession,
+ now: Optional[datetime] = None,
+ limit: int = 100,
+) -> List[ActiveDiscount]:
+ """Get expired active discount reservations for cleanup/notifications."""
+ now_utc = now or datetime.now(timezone.utc)
+ stmt = (
+ select(ActiveDiscount)
+ .where(ActiveDiscount.expires_at <= now_utc)
+ .order_by(ActiveDiscount.expires_at.asc())
+ .limit(limit)
+ )
+ result = await session.execute(stmt)
+ return list(result.scalars().all())
+
+
async def clear_active_discounts_by_promo_code(
session: AsyncSession,
promo_code_id: int
diff --git a/db/models.py b/db/models.py
index cc55132..61db14a 100644
--- a/db/models.py
+++ b/db/models.py
@@ -207,7 +207,7 @@ class PromoCodeActivation(Base):
class ActiveDiscount(Base):
- """Tracks pending discount promo codes awaiting payment (permanent until used)"""
+ """Tracks pending discount promo code reservations awaiting payment."""
__tablename__ = "active_discounts"
user_id = Column(
@@ -222,6 +222,7 @@ class ActiveDiscount(Base):
)
discount_percentage = Column(Integer, nullable=False)
activated_at = Column(DateTime(timezone=True), server_default=func.now())
+ expires_at = Column(DateTime(timezone=True), nullable=False)
promo_code = relationship("PromoCode")
user = relationship("User")
diff --git a/locales/en.json b/locales/en.json
index 306098e..a926ddd 100644
--- a/locales/en.json
+++ b/locales/en.json
@@ -78,6 +78,7 @@
"promo_code_applied_success_full": "✅ Promo code applied successfully!\nSubscription active until {end_date}.\n\nConnection key:\n{config_link}\n\nTo connect, open the link and follow the instructions 👇",
"discount_promo_code_applied_success": "✅ Promo code {code} activated!\n\n💰 A {discount}% discount will be applied to your next purchase.\n\nSelect a plan for payment.",
"discount_promo_already_active": "❌ You already have an active discount promo code ({code}, -{discount_pct}%). Use it first or wait until the payment is complete.",
+ "discount_promo_expired_need_reactivate": "⏰ The discount hold period has ended{code_part}.\n\nPlease enter the promo code again to apply the discount.",
"promo_code_not_found_or_not_discount": "❌ Promo code {code} not found or is invalid.",
"active_discount_notice": "🎁 Active discount: {code} (-{discount_pct}%)\n💵 Price: {original_price}{currency_symbol} ➔ {discounted_price}{currency_symbol}\n💰 Savings: {discount_amount}{currency_symbol}",
"error_applying_promo_bonus": "Failed to apply promo bonus. Please try again later or contact support.",
diff --git a/locales/ru.json b/locales/ru.json
index 7913ce2..5383636 100644
--- a/locales/ru.json
+++ b/locales/ru.json
@@ -78,6 +78,7 @@
"promo_code_applied_success_full": "✅ Промокод успешно применен!\nПодписка активна до {end_date}.\n\nКлюч подключения:\n{config_link}\n\nЧтобы подключиться, перейдите по ссылке и следуйте инструкции 👇",
"discount_promo_code_applied_success": "✅ Промокод {code} активирован!\n\n💰 Скидка {discount}% будет применена к вашей следующей покупке.\n\nВыберите тариф для оплаты.",
"discount_promo_already_active": "❌ У вас уже есть активированный промокод на скидку ({code}, -{discount_pct}%). Используйте его сначала или дождитесь окончания платежа.",
+ "discount_promo_expired_need_reactivate": "⏰ Время действия скидки истекло{code_part}.\n\nВведите промокод снова, чтобы применить скидку.",
"promo_code_not_found_or_not_discount": "❌ Промокод {code} не найден или недействителен.",
"active_discount_notice": "🎁 Активна скидка: {code} (-{discount_pct}%)\n💵 Цена: {original_price}{currency_symbol} ➔ {discounted_price}{currency_symbol}\n💰 Экономия: {discount_amount}{currency_symbol}",
"error_applying_promo_bonus": "Не удалось применить бонус по промокоду. Пожалуйста, попробуйте позже или свяжитесь с поддержкой.",