fix: reset trial eligibility from web admin

This commit is contained in:
3252a8
2026-06-01 19:30:16 +03:00
parent 74272039c5
commit 687fc03e8c
16 changed files with 253 additions and 32 deletions
+6 -9
View File
@@ -1172,10 +1172,6 @@ async def admin_user_reset_trial_route(request: web.Request) -> web.Response:
actor_id = _require_admin_user_id(request) actor_id = _require_admin_user_id(request)
target_id = int(request.match_info["user_id"]) target_id = int(request.match_info["user_id"])
settings: Settings = request.app["settings"] settings: Settings = request.app["settings"]
panel_service = request.app.get("panel_service")
subscription_service = request.app.get("subscription_service")
if panel_service is None or subscription_service is None:
return _error(503, "service_unavailable")
async_session_factory: sessionmaker = request.app["async_session_factory"] async_session_factory: sessionmaker = request.app["async_session_factory"]
async with async_session_factory() as session: async with async_session_factory() as session:
@@ -1183,16 +1179,17 @@ async def admin_user_reset_trial_route(request: web.Request) -> web.Response:
if not user: if not user:
return _error(404, "not_found") return _error(404, "not_found")
active = await subscription_dal.get_active_subscription_by_user_id(session, target_id) reset_at = await user_dal.mark_trial_eligibility_reset(session, target_id)
if active: if reset_at is None:
await session.delete(active) await session.rollback()
return _error(404, "not_found")
await message_log_dal.create_message_log( await message_log_dal.create_message_log_no_commit(
session, session,
{ {
"user_id": actor_id, "user_id": actor_id,
"event_type": "admin_reset_trial_webapp", "event_type": "admin_reset_trial_webapp",
"content": f"Reset trial for user_id={target_id}", "content": f"Reset trial eligibility for user_id={target_id}",
"is_admin_event": True, "is_admin_event": True,
"target_user_id": target_id, "target_user_id": target_id,
}, },
+1 -1
View File
@@ -67,7 +67,7 @@ async def _build_user_payload(request: web.Request, user_id: int) -> Dict[str, A
trial_available = bool( trial_available = bool(
settings.TRIAL_ENABLED settings.TRIAL_ENABLED
and settings.TRIAL_DURATION_DAYS > 0 and settings.TRIAL_DURATION_DAYS > 0
and not await subscription_service.has_had_any_subscription(session, user_id) and not await subscription_service.has_trial_blocking_subscription(session, user_id)
) )
avatar = await _ensure_cached_telegram_avatar(request, session, db_user) avatar = await _ensure_cached_telegram_avatar(request, session, db_user)
try: try:
@@ -932,8 +932,7 @@ async def handle_reset_trial(
_ = lambda key, **kwargs: i18n_instance.gettext(lang, key, **kwargs) _ = lambda key, **kwargs: i18n_instance.gettext(lang, key, **kwargs)
try: try:
# Delete all user subscriptions to reset trial eligibility await user_dal.mark_trial_eligibility_reset(session, user.user_id)
await subscription_dal.delete_all_user_subscriptions(session, user.user_id)
await session.commit() await session.commit()
await callback.answer(_("admin_user_trial_reset_success"), show_alert=True) await callback.answer(_("admin_user_trial_reset_success"), show_alert=True)
+4 -4
View File
@@ -45,12 +45,12 @@ async def should_show_trial_button(
if not settings.TRIAL_ENABLED: if not settings.TRIAL_ENABLED:
return False return False
if hasattr(subscription_service, "has_had_any_subscription") and callable( if hasattr(subscription_service, "has_trial_blocking_subscription") and callable(
getattr(subscription_service, "has_had_any_subscription") getattr(subscription_service, "has_trial_blocking_subscription")
): ):
return not await subscription_service.has_had_any_subscription(session, user_id) return not await subscription_service.has_trial_blocking_subscription(session, user_id)
logging.error("Method has_had_any_subscription is missing in SubscriptionService!") logging.error("Method has_trial_blocking_subscription is missing in SubscriptionService!")
return False return False
+9 -7
View File
@@ -46,7 +46,7 @@ async def request_trial_confirmation_handler(
return return
if settings.TRIAL_ENABLED: if settings.TRIAL_ENABLED:
if not await subscription_service.has_had_any_subscription(session, user_id): if not await subscription_service.has_trial_blocking_subscription(session, user_id):
pass pass
if not settings.TRIAL_ENABLED: if not settings.TRIAL_ENABLED:
@@ -60,7 +60,7 @@ async def request_trial_confirmation_handler(
pass pass
return return
if await subscription_service.has_had_any_subscription(session, user_id): if await subscription_service.has_trial_blocking_subscription(session, user_id):
await callback.message.edit_text( await callback.message.edit_text(
_("trial_already_had_subscription_or_trial"), _("trial_already_had_subscription_or_trial"),
reply_markup=get_main_menu_inline_keyboard(current_lang, i18n, settings, False), reply_markup=get_main_menu_inline_keyboard(current_lang, i18n, settings, False),
@@ -147,8 +147,9 @@ async def request_trial_confirmation_handler(
await callback.answer(final_message_text_in_chat, show_alert=True) await callback.answer(final_message_text_in_chat, show_alert=True)
except Exception: except Exception:
pass pass
if settings.TRIAL_ENABLED and not await subscription_service.has_had_any_subscription( if (
session, user_id settings.TRIAL_ENABLED
and not await subscription_service.has_trial_blocking_subscription(session, user_id)
): ):
show_trial_button_after_action = True show_trial_button_after_action = True
@@ -218,7 +219,7 @@ async def confirm_activate_trial_handler(
callback, settings, i18n_data, subscription_service, session, is_edit=True callback, settings, i18n_data, subscription_service, session, is_edit=True
) )
return return
if await subscription_service.has_had_any_subscription(session, user_id): if await subscription_service.has_trial_blocking_subscription(session, user_id):
try: try:
await callback.answer(_("trial_already_had_subscription_or_trial"), show_alert=True) await callback.answer(_("trial_already_had_subscription_or_trial"), show_alert=True)
except Exception: except Exception:
@@ -283,8 +284,9 @@ async def confirm_activate_trial_handler(
await callback.answer(final_message_text_in_chat, show_alert=True) await callback.answer(final_message_text_in_chat, show_alert=True)
except Exception: except Exception:
pass pass
if settings.TRIAL_ENABLED and not await subscription_service.has_had_any_subscription( if (
session, user_id settings.TRIAL_ENABLED
and not await subscription_service.has_trial_blocking_subscription(session, user_id)
): ):
show_trial_button_after_action = True show_trial_button_after_action = True
@@ -62,6 +62,9 @@ class PaymentContextMixin:
async def has_had_any_subscription(self, session: AsyncSession, user_id: int) -> bool: async def has_had_any_subscription(self, session: AsyncSession, user_id: int) -> bool:
return await subscription_dal.has_any_subscription_for_user(session, user_id) return await subscription_dal.has_any_subscription_for_user(session, user_id)
async def has_trial_blocking_subscription(self, session: AsyncSession, user_id: int) -> bool:
return await subscription_dal.has_trial_blocking_subscription_for_user(session, user_id)
async def has_active_subscription(self, session: AsyncSession, user_id: int) -> bool: async def has_active_subscription(self, session: AsyncSession, user_id: int) -> bool:
"""Return True if user currently has an active subscription (end_date in future).""" """Return True if user currently has an active subscription (end_date in future)."""
try: try:
@@ -22,7 +22,7 @@ class TrialSubscriptionMixin:
"message_key": "user_not_found_for_trial", "message_key": "user_not_found_for_trial",
} }
if await self.has_had_any_subscription(session, user_id): if await self.has_trial_blocking_subscription(session, user_id):
return { return {
"eligible": False, "eligible": False,
"activated": False, "activated": False,
+25 -3
View File
@@ -4,12 +4,12 @@ import secrets
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
from sqlalchemy import delete, func, or_, update from sqlalchemy import and_, delete, func, or_, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.future import select from sqlalchemy.future import select
from sqlalchemy.orm import selectinload from sqlalchemy.orm import selectinload
from db.models import Subscription, SubscriptionNotification from db.models import Subscription, SubscriptionNotification, User
INSTALL_SHARE_TOKEN_BYTES = 16 INSTALL_SHARE_TOKEN_BYTES = 16
@@ -252,7 +252,7 @@ async def deactivate_all_user_subscriptions(session: AsyncSession, user_id: int)
async def delete_all_user_subscriptions(session: AsyncSession, user_id: int) -> int: async def delete_all_user_subscriptions(session: AsyncSession, user_id: int) -> int:
"""Completely delete all user subscriptions (for trial reset)""" """Completely delete all user subscriptions."""
stmt = delete(Subscription).where(Subscription.user_id == user_id) stmt = delete(Subscription).where(Subscription.user_id == user_id)
result = await session.execute(stmt) result = await session.execute(stmt)
if result.rowcount > 0: if result.rowcount > 0:
@@ -284,6 +284,28 @@ async def has_any_subscription_for_user(session: AsyncSession, user_id: int) ->
return result.scalar_one_or_none() is not None return result.scalar_one_or_none() is not None
async def has_trial_blocking_subscription_for_user(session: AsyncSession, user_id: int) -> bool:
now_utc = datetime.now(timezone.utc)
reset_at = (
select(User.trial_eligibility_reset_at).where(User.user_id == user_id).scalar_subquery()
)
subscription_anchor = func.coalesce(Subscription.start_date, Subscription.end_date)
stmt = (
select(Subscription.subscription_id)
.where(
Subscription.user_id == user_id,
or_(
reset_at.is_(None),
and_(Subscription.is_active == True, Subscription.end_date > now_utc),
subscription_anchor > reset_at,
),
)
.limit(1)
)
result = await session.execute(stmt)
return result.scalar_one_or_none() is not None
async def get_subscriptions_near_expiration( async def get_subscriptions_near_expiration(
session: AsyncSession, days_threshold: int session: AsyncSession, days_threshold: int
) -> List[Subscription]: ) -> List[Subscription]:
+14
View File
@@ -261,6 +261,20 @@ async def create_email_user(
) )
async def mark_trial_eligibility_reset(
session: AsyncSession,
user_id: int,
*,
reset_at: Optional[datetime] = None,
) -> Optional[datetime]:
reset_at = reset_at or datetime.now(timezone.utc)
stmt = update(User).where(User.user_id == user_id).values(trial_eligibility_reset_at=reset_at)
result = await session.execute(stmt)
if result.rowcount <= 0:
return None
return reset_at
async def _has_active_panel_subscription( async def _has_active_panel_subscription(
session: AsyncSession, user_id: int, panel_user_uuid: str session: AsyncSession, user_id: int, panel_user_uuid: str
) -> bool: ) -> bool:
+14
View File
@@ -1061,6 +1061,15 @@ def _migration_0032_add_telegram_notification_status(connection: Connection) ->
connection.execute(text(f"ALTER TABLE users ADD COLUMN {column} {ddl_type}")) connection.execute(text(f"ALTER TABLE users ADD COLUMN {column} {ddl_type}"))
def _migration_0033_add_trial_eligibility_reset_marker(connection: Connection) -> None:
inspector = inspect(connection)
columns: Set[str] = {col["name"] for col in inspector.get_columns("users")}
if "trial_eligibility_reset_at" not in columns:
connection.execute(
text("ALTER TABLE users ADD COLUMN trial_eligibility_reset_at TIMESTAMPTZ")
)
MIGRATIONS: List[Migration] = [ MIGRATIONS: List[Migration] = [
Migration( Migration(
id="0001_add_channel_subscription_fields", id="0001_add_channel_subscription_fields",
@@ -1233,6 +1242,11 @@ MIGRATIONS: List[Migration] = [
description="Track whether the bot can message Telegram-linked users", description="Track whether the bot can message Telegram-linked users",
upgrade=_migration_0032_add_telegram_notification_status, upgrade=_migration_0032_add_telegram_notification_status,
), ),
Migration(
id="0033_add_trial_eligibility_reset_marker",
description="Track admin resets of per-user trial eligibility without deleting history",
upgrade=_migration_0033_add_trial_eligibility_reset_marker,
),
] ]
+1
View File
@@ -47,6 +47,7 @@ class User(Base):
referred_by_id = Column(BigInteger, ForeignKey("users.user_id"), nullable=True) referred_by_id = Column(BigInteger, ForeignKey("users.user_id"), nullable=True)
lifetime_used_traffic_bytes = Column(BigInteger, nullable=True) lifetime_used_traffic_bytes = Column(BigInteger, nullable=True)
lifetime_used_traffic_synced_at = Column(DateTime(timezone=True), nullable=True) lifetime_used_traffic_synced_at = Column(DateTime(timezone=True), nullable=True)
trial_eligibility_reset_at = Column(DateTime(timezone=True), nullable=True)
channel_subscription_verified = Column(Boolean, nullable=True) channel_subscription_verified = Column(Boolean, nullable=True)
channel_subscription_checked_at = Column(DateTime(timezone=True), nullable=True) channel_subscription_checked_at = Column(DateTime(timezone=True), nullable=True)
channel_subscription_verified_for = Column(BigInteger, nullable=True) channel_subscription_verified_for = Column(BigInteger, nullable=True)
+5 -2
View File
@@ -469,8 +469,11 @@ export function createUsersStore({ api, onToast, at, routePrefix = "" }) {
state.update((st) => ({ ...st, userActionBusy: true })); state.update((st) => ({ ...st, userActionBusy: true }));
try { try {
const res = await api(`/admin/users/${s.openedUser.user_id}/reset-trial`, { method: "POST" }); const res = await api(`/admin/users/${s.openedUser.user_id}/reset-trial`, { method: "POST" });
if (res?.ok) onToast(at("trial_reset", {}, "Триал сброшен")); if (res?.ok) {
else onToast(res?.error || at("error", {}, "Ошибка")); onToast(at("trial_reset", {}, "Триал сброшен"));
await openUser(s.openedUser.user_id, { skipPush: true, pathContext: _pathContext });
if (_activeRef === "users") await loadUsers();
} else onToast(res?.error || at("error", {}, "Ошибка"));
} finally { } finally {
state.update((st) => ({ ...st, userActionBusy: false })); state.update((st) => ({ ...st, userActionBusy: false }));
} }
+77
View File
@@ -0,0 +1,77 @@
import json
import unittest
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
from bot.app.web.admin_api_impl import users as admin_users
class FakeSession:
def __init__(self):
self.committed = False
self.rolled_back = False
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, exc, tb):
return False
async def commit(self):
self.committed = True
async def rollback(self):
self.rolled_back = True
class AdminUserResetTrialRouteTests(unittest.IsolatedAsyncioTestCase):
def _request(self, session: FakeSession):
return SimpleNamespace(
app={
"settings": SimpleNamespace(),
"async_session_factory": lambda: session,
},
match_info={"user_id": "42"},
)
async def test_marks_trial_reset_without_deleting_subscription_history(self):
session = FakeSession()
request = self._request(session)
user = SimpleNamespace(user_id=42)
with (
patch.object(admin_users, "_require_admin_user_id", return_value=100),
patch.object(admin_users.user_dal, "get_user_by_id", AsyncMock(return_value=user)),
patch.object(
admin_users.user_dal,
"mark_trial_eligibility_reset",
AsyncMock(return_value=object()),
) as mark_reset,
patch.object(
admin_users.subscription_dal,
"delete_all_user_subscriptions",
AsyncMock(),
) as delete_all,
patch.object(
admin_users.message_log_dal, "create_message_log_no_commit", AsyncMock()
) as log,
patch.object(
admin_users, "_invalidate_after_admin_user_mutation", AsyncMock()
) as invalidate,
):
response = await admin_users.admin_user_reset_trial_route(request)
self.assertEqual(response.status, 200)
self.assertEqual(json.loads(response.text)["ok"], True)
mark_reset.assert_awaited_once_with(session, 42)
delete_all.assert_not_awaited()
log_payload = log.await_args.args[1]
self.assertEqual(log_payload["event_type"], "admin_reset_trial_webapp")
self.assertEqual(log_payload["target_user_id"], 42)
invalidate.assert_awaited_once()
self.assertTrue(session.committed)
self.assertFalse(session.rolled_back)
if __name__ == "__main__":
unittest.main()
+43 -2
View File
@@ -202,7 +202,7 @@ class SubscriptionServiceActivationDispatchTests(unittest.IsolatedAsyncioTestCas
TRIAL_SQUAD_UUIDS="trial-squad", TRIAL_SQUAD_UUIDS="trial-squad",
) )
service = _make_service(settings) service = _make_service(settings)
service.has_had_any_subscription = AsyncMock(return_value=False) service.has_trial_blocking_subscription = AsyncMock(return_value=False)
service._get_or_create_panel_user_link_details = AsyncMock( service._get_or_create_panel_user_link_details = AsyncMock(
return_value=("panel-user", "panel-sub", "short", True) return_value=("panel-user", "panel-sub", "short", True)
) )
@@ -255,7 +255,7 @@ class SubscriptionServiceActivationDispatchTests(unittest.IsolatedAsyncioTestCas
TRIAL_SQUAD_UUIDS=" , ", TRIAL_SQUAD_UUIDS=" , ",
) )
service = _make_service(settings) service = _make_service(settings)
service.has_had_any_subscription = AsyncMock(return_value=False) service.has_trial_blocking_subscription = AsyncMock(return_value=False)
service._get_or_create_panel_user_link_details = AsyncMock( service._get_or_create_panel_user_link_details = AsyncMock(
return_value=("panel-user", "panel-sub", "short", True) return_value=("panel-user", "panel-sub", "short", True)
) )
@@ -292,6 +292,47 @@ class SubscriptionServiceActivationDispatchTests(unittest.IsolatedAsyncioTestCas
panel_payload = service.panel_service.update_user_details_on_panel.await_args.args[1] panel_payload = service.panel_service.update_user_details_on_panel.await_args.args[1]
self.assertEqual(panel_payload["activeInternalSquads"], ["fallback-a", "fallback-b"]) self.assertEqual(panel_payload["activeInternalSquads"], ["fallback-a", "fallback-b"])
async def test_activate_trial_rejects_users_with_blocking_subscription_history(self):
with tempfile.TemporaryDirectory() as tmpdir:
settings = _make_settings(
_tariffs_config_payload(),
tmpdir,
TRIAL_ENABLED=True,
TRIAL_DURATION_DAYS=3,
)
service = _make_service(settings)
service.has_trial_blocking_subscription = AsyncMock(return_value=True)
service._get_or_create_panel_user_link_details = AsyncMock()
service.panel_service.update_user_details_on_panel = AsyncMock()
session = AsyncMock()
db_user = SimpleNamespace(
user_id=42,
telegram_id=42,
email=None,
username="trial-user",
first_name="Trial",
last_name="User",
)
with (
patch(
"bot.services.subscription_service_impl.trial.user_dal.get_user_by_id",
AsyncMock(return_value=db_user),
),
patch(
"bot.services.subscription_service_impl.trial.subscription_dal.upsert_subscription",
AsyncMock(),
) as upsert_subscription,
):
result = await service.activate_trial_subscription(session, user_id=42)
self.assertFalse(result["activated"])
self.assertFalse(result["eligible"])
self.assertEqual(result["message_key"], "trial_already_had_subscription_or_trial")
service._get_or_create_panel_user_link_details.assert_not_awaited()
service.panel_service.update_user_details_on_panel.assert_not_awaited()
upsert_subscription.assert_not_awaited()
async def test_activate_subscription_dispatches_traffic_sale_mode(self): async def test_activate_subscription_dispatches_traffic_sale_mode(self):
with tempfile.TemporaryDirectory() as tmpdir: with tempfile.TemporaryDirectory() as tmpdir:
settings = _make_settings(_tariffs_config_payload(), tmpdir) settings = _make_settings(_tariffs_config_payload(), tmpdir)
+10
View File
@@ -29,3 +29,13 @@ def test_support_models_expose_expected_tables():
def test_user_model_tracks_lifetime_traffic_sync_timestamp(): def test_user_model_tracks_lifetime_traffic_sync_timestamp():
assert "lifetime_used_traffic_synced_at" in User.__table__.columns assert "lifetime_used_traffic_synced_at" in User.__table__.columns
def test_trial_eligibility_reset_migration_and_model_are_registered():
ids = [migration.id for migration in MIGRATIONS]
assert "0033_add_trial_eligibility_reset_marker" in ids
assert ids.index("0033_add_trial_eligibility_reset_marker") > ids.index(
"0032_add_telegram_notification_status"
)
assert "trial_eligibility_reset_at" in User.__table__.columns
+39 -1
View File
@@ -6,7 +6,7 @@ from unittest.mock import AsyncMock, patch
from sqlalchemy.dialects import postgresql from sqlalchemy.dialects import postgresql
from sqlalchemy.sql.dml import Delete, Update from sqlalchemy.sql.dml import Delete, Update
from db.dal import user_dal from db.dal import subscription_dal, user_dal
class FakeResult: class FakeResult:
@@ -122,6 +122,44 @@ class UserDalReferralTests(unittest.IsolatedAsyncioTestCase):
self.assertIs(result, referrer) self.assertIs(result, referrer)
async def test_mark_trial_eligibility_reset_updates_user_marker(self):
reset_at = datetime(2026, 6, 1, tzinfo=timezone.utc)
session = SimpleNamespace(execute=AsyncMock(return_value=FakeResult(rowcount=1)))
result = await user_dal.mark_trial_eligibility_reset(session, 42, reset_at=reset_at)
self.assertEqual(result, reset_at)
stmt = session.execute.await_args.args[0]
sql = str(
stmt.compile(
dialect=postgresql.dialect(),
compile_kwargs={"literal_binds": True},
)
).upper()
self.assertIn("UPDATE USERS", sql)
self.assertIn("TRIAL_ELIGIBILITY_RESET_AT", sql)
self.assertIn("USER_ID = 42", sql)
class SubscriptionDalTrialEligibilityTests(unittest.IsolatedAsyncioTestCase):
async def test_trial_blocking_history_honors_user_reset_marker(self):
session = SimpleNamespace(execute=AsyncMock(return_value=FakeResult(7)))
result = await subscription_dal.has_trial_blocking_subscription_for_user(session, 42)
self.assertTrue(result)
stmt = session.execute.await_args.args[0]
sql = str(
stmt.compile(
dialect=postgresql.dialect(),
compile_kwargs={"literal_binds": True},
)
).upper()
self.assertIn("TRIAL_ELIGIBILITY_RESET_AT", sql)
self.assertIn("SUBSCRIPTIONS.IS_ACTIVE = TRUE", sql)
self.assertIn("COALESCE(SUBSCRIPTIONS.START_DATE, SUBSCRIPTIONS.END_DATE)", sql)
self.assertIn("SUBSCRIPTIONS.USER_ID = 42", sql)
class UserDalMergeTests(unittest.IsolatedAsyncioTestCase): class UserDalMergeTests(unittest.IsolatedAsyncioTestCase):
async def test_get_panel_user_uuids_for_user_includes_subscription_fallbacks_once(self): async def test_get_panel_user_uuids_for_user_includes_subscription_fallbacks_once(self):