Merge pull request #2390 from Gy9vin/main

fix(referral): исправить потерю реферальных кодов при обязательной по…
This commit is contained in:
Egor
2026-01-22 22:00:27 +03:00
committed by GitHub
4 changed files with 389 additions and 26 deletions
+47 -19
View File
@@ -48,6 +48,10 @@ from app.utils.promo_offer import (
)
from app.utils.timezone import format_local_datetime
from app.database.crud.user_message import get_random_active_message
from app.middlewares.channel_checker import (
get_pending_payload_from_redis,
delete_pending_payload_from_redis,
)
from app.database.crud.subscription import decrement_subscription_server_counts
from app.services.blacklist_service import blacklist_service
@@ -327,6 +331,19 @@ async def cmd_start(message: types.Message, state: FSMContext, db: AsyncSession,
campaign_notification_sent = data.pop("campaign_notification_sent", False)
state_needs_update = had_pending_payload or had_campaign_notification_flag
# Если в FSM state нет payload, пробуем получить из Redis (резервный механизм)
if not pending_start_payload:
redis_payload = await get_pending_payload_from_redis(message.from_user.id)
if redis_payload:
pending_start_payload = redis_payload
state_needs_update = True
logger.info(
"📦 START: Payload '%s' восстановлен из Redis (fallback)",
pending_start_payload,
)
# Очищаем Redis после получения
await delete_pending_payload_from_redis(message.from_user.id)
referral_code = None
campaign = None
start_args = message.text.split()
@@ -1832,6 +1849,17 @@ async def required_sub_channel_check(
state_data = await state.get_data() or {}
pending_start_payload = state_data.pop("pending_start_payload", None)
# Если в FSM state нет payload, пробуем получить из Redis (резервный механизм)
if not pending_start_payload:
redis_payload = await get_pending_payload_from_redis(query.from_user.id)
if redis_payload:
pending_start_payload = redis_payload
logger.info(
"📦 CHANNEL CHECK: Payload '%s' восстановлен из Redis (fallback)",
pending_start_payload,
)
state_updated = pending_start_payload is not None
if pending_start_payload:
@@ -1840,27 +1868,27 @@ async def required_sub_channel_check(
pending_start_payload,
)
if "campaign_id" not in state_data and "referral_code" not in state_data:
campaign = await get_campaign_by_start_parameter(
db,
pending_start_payload,
only_active=True,
)
# Очищаем Redis после получения payload
await delete_pending_payload_from_redis(query.from_user.id)
if campaign:
state_data["campaign_id"] = campaign.id
logger.info(
"📣 CHANNEL CHECK: Кампания %s восстановлена из payload",
campaign.id,
)
else:
state_data["referral_code"] = pending_start_payload
logger.info(
"🎯 CHANNEL CHECK: Payload интерпретирован как реферальный код",
)
# Всегда обновляем referral_code если есть новый payload
# (исправление бага с устаревшими данными в state)
campaign = await get_campaign_by_start_parameter(
db,
pending_start_payload,
only_active=True,
)
if campaign:
state_data["campaign_id"] = campaign.id
logger.info(
"📣 CHANNEL CHECK: Кампания %s восстановлена из payload",
campaign.id,
)
else:
logger.debug(
"️ CHANNEL CHECK: Payload уже обработан ранее, пропускаем восстановление",
state_data["referral_code"] = pending_start_payload
logger.info(
"🎯 CHANNEL CHECK: Payload интерпретирован как реферальный код",
)
if state_updated:
+78 -7
View File
@@ -7,6 +7,7 @@ from aiogram.fsm.context import FSMContext
from aiogram.types import TelegramObject, Update, Message, CallbackQuery
from aiogram.enums import ChatMemberStatus
from sqlalchemy.ext.asyncio import AsyncSession
import redis.asyncio as aioredis
from app.config import settings
from app.database.database import AsyncSessionLocal
@@ -23,6 +24,58 @@ from app.services.admin_notification_service import AdminNotificationService
logger = logging.getLogger(__name__)
# Ключ для хранения pending_start_payload в Redis (резервный механизм)
REDIS_PAYLOAD_KEY_PREFIX = "pending_start_payload:"
REDIS_PAYLOAD_TTL = 3600 # 1 час
async def save_pending_payload_to_redis(telegram_id: int, payload: str) -> bool:
"""Сохраняет pending_start_payload в Redis напрямую (резервный механизм)."""
try:
redis_client = aioredis.from_url(settings.REDIS_URL)
key = f"{REDIS_PAYLOAD_KEY_PREFIX}{telegram_id}"
await redis_client.set(key, payload, ex=REDIS_PAYLOAD_TTL)
await redis_client.aclose()
logger.info(
"💾 [Redis fallback] Сохранен payload '%s' для пользователя %s",
payload,
telegram_id,
)
return True
except Exception as e:
logger.error(
"❌ [Redis fallback] Ошибка сохранения payload для %s: %s",
telegram_id,
e,
)
return False
async def get_pending_payload_from_redis(telegram_id: int) -> Optional[str]:
"""Получает pending_start_payload из Redis (резервный механизм)."""
try:
redis_client = aioredis.from_url(settings.REDIS_URL)
key = f"{REDIS_PAYLOAD_KEY_PREFIX}{telegram_id}"
payload = await redis_client.get(key)
await redis_client.aclose()
if payload:
return payload.decode("utf-8") if isinstance(payload, bytes) else payload
return None
except Exception as e:
logger.debug("❌ [Redis fallback] Ошибка получения payload для %s: %s", telegram_id, e)
return None
async def delete_pending_payload_from_redis(telegram_id: int) -> None:
"""Удаляет pending_start_payload из Redis."""
try:
redis_client = aioredis.from_url(settings.REDIS_URL)
key = f"{REDIS_PAYLOAD_KEY_PREFIX}{telegram_id}"
await redis_client.delete(key)
await redis_client.aclose()
except Exception:
pass
class ChannelCheckerMiddleware(BaseMiddleware):
"""
@@ -170,8 +223,11 @@ class ChannelCheckerMiddleware(BaseMiddleware):
event: TelegramObject,
bot: Optional[Bot] = None,
) -> None:
if not state:
return
telegram_id = None
if isinstance(event, Message):
telegram_id = event.from_user.id if event.from_user else None
elif isinstance(event, CallbackQuery):
telegram_id = event.from_user.id if event.from_user else None
message: Optional[Message] = None
if isinstance(event, Message):
@@ -194,11 +250,26 @@ class ChannelCheckerMiddleware(BaseMiddleware):
payload = parts[1]
state_data = await state.get_data() or {}
if state_data.get("pending_start_payload") != payload:
state_data["pending_start_payload"] = payload
await state.set_data(state_data)
logger.debug("💾 Сохранен start payload %s для последующей обработки", payload)
# Сохраняем в FSM state
if state:
state_data = await state.get_data() or {}
if state_data.get("pending_start_payload") != payload:
state_data["pending_start_payload"] = payload
await state.set_data(state_data)
logger.info(
"💾 Сохранен start payload '%s' для пользователя %s (FSM)",
payload,
telegram_id,
)
else:
logger.warning(
"⚠️ _capture_start_payload: state=None для пользователя %s",
telegram_id,
)
# Также сохраняем в Redis как резерв (на случай потери FSM state)
if telegram_id:
await save_pending_payload_to_redis(telegram_id, payload)
if bot and message.from_user:
await self._try_send_campaign_visit_notification(
+1
View File
@@ -0,0 +1 @@
# Middlewares tests package
@@ -0,0 +1,263 @@
"""Тесты для функций сохранения/получения pending_start_payload в channel_checker."""
from pathlib import Path
import sys
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch, create_autospec
import pytest
ROOT_DIR = Path(__file__).resolve().parents[2]
if str(ROOT_DIR) not in sys.path:
sys.path.insert(0, str(ROOT_DIR))
from aiogram.types import Message, User
class TestRedisPayloadFunctions:
"""Тесты для Redis-функций сохранения payload."""
async def test_save_pending_payload_to_redis_success(self, monkeypatch):
"""Тест успешного сохранения payload в Redis."""
from app.middlewares import channel_checker
mock_redis = AsyncMock()
mock_redis.set = AsyncMock(return_value=True)
mock_redis.aclose = AsyncMock()
with patch("app.middlewares.channel_checker.aioredis") as mock_aioredis:
mock_aioredis.from_url = MagicMock(return_value=mock_redis)
result = await channel_checker.save_pending_payload_to_redis(123456, "ref_test123")
assert result is True
mock_redis.set.assert_awaited_once()
call_args = mock_redis.set.await_args
assert "pending_start_payload:123456" in call_args.args[0]
assert call_args.args[1] == "ref_test123"
assert call_args.kwargs.get("ex") == 3600
mock_redis.aclose.assert_awaited_once()
async def test_save_pending_payload_to_redis_failure(self, monkeypatch):
"""Тест обработки ошибки при сохранении в Redis."""
from app.middlewares import channel_checker
with patch("app.middlewares.channel_checker.aioredis") as mock_aioredis:
mock_aioredis.from_url = MagicMock(side_effect=Exception("Redis connection failed"))
result = await channel_checker.save_pending_payload_to_redis(123456, "ref_test123")
assert result is False
async def test_get_pending_payload_from_redis_success(self, monkeypatch):
"""Тест успешного получения payload из Redis."""
from app.middlewares import channel_checker
mock_redis = AsyncMock()
mock_redis.get = AsyncMock(return_value=b"ref_test123")
mock_redis.aclose = AsyncMock()
with patch("app.middlewares.channel_checker.aioredis") as mock_aioredis:
mock_aioredis.from_url = MagicMock(return_value=mock_redis)
result = await channel_checker.get_pending_payload_from_redis(123456)
assert result == "ref_test123"
mock_redis.get.assert_awaited_once()
mock_redis.aclose.assert_awaited_once()
async def test_get_pending_payload_from_redis_not_found(self, monkeypatch):
"""Тест когда payload не найден в Redis."""
from app.middlewares import channel_checker
mock_redis = AsyncMock()
mock_redis.get = AsyncMock(return_value=None)
mock_redis.aclose = AsyncMock()
with patch("app.middlewares.channel_checker.aioredis") as mock_aioredis:
mock_aioredis.from_url = MagicMock(return_value=mock_redis)
result = await channel_checker.get_pending_payload_from_redis(123456)
assert result is None
async def test_get_pending_payload_from_redis_failure(self, monkeypatch):
"""Тест обработки ошибки при получении из Redis."""
from app.middlewares import channel_checker
with patch("app.middlewares.channel_checker.aioredis") as mock_aioredis:
mock_aioredis.from_url = MagicMock(side_effect=Exception("Redis connection failed"))
result = await channel_checker.get_pending_payload_from_redis(123456)
assert result is None
async def test_delete_pending_payload_from_redis(self, monkeypatch):
"""Тест удаления payload из Redis."""
from app.middlewares import channel_checker
mock_redis = AsyncMock()
mock_redis.delete = AsyncMock(return_value=1)
mock_redis.aclose = AsyncMock()
with patch("app.middlewares.channel_checker.aioredis") as mock_aioredis:
mock_aioredis.from_url = MagicMock(return_value=mock_redis)
# Не должно бросать исключение
await channel_checker.delete_pending_payload_from_redis(123456)
mock_redis.delete.assert_awaited_once()
async def test_delete_pending_payload_from_redis_handles_error(self, monkeypatch):
"""Тест что удаление не бросает исключение при ошибке."""
from app.middlewares import channel_checker
with patch("app.middlewares.channel_checker.aioredis") as mock_aioredis:
mock_aioredis.from_url = MagicMock(side_effect=Exception("Redis error"))
# Не должно бросать исключение
await channel_checker.delete_pending_payload_from_redis(123456)
def _create_mock_message(text: str, user_id: int):
"""Создаёт мок Message с нужными атрибутами."""
mock_msg = MagicMock(spec=Message)
mock_msg.text = text
mock_msg.from_user = SimpleNamespace(id=user_id)
return mock_msg
class TestCaptureStartPayload:
"""Тесты для метода _capture_start_payload."""
async def test_capture_saves_to_fsm_state(self, monkeypatch):
"""Тест сохранения payload в FSM state."""
from app.middlewares.channel_checker import ChannelCheckerMiddleware
middleware = ChannelCheckerMiddleware()
mock_state = AsyncMock()
mock_state.get_data = AsyncMock(return_value={})
mock_state.set_data = AsyncMock()
mock_message = _create_mock_message("/start ref_abc123", 123456)
with patch("app.middlewares.channel_checker.save_pending_payload_to_redis", new_callable=AsyncMock) as mock_save_redis:
await middleware._capture_start_payload(mock_state, mock_message, None)
mock_state.set_data.assert_awaited_once()
saved_data = mock_state.set_data.await_args.args[0]
assert saved_data["pending_start_payload"] == "ref_abc123"
# Также должен сохраняться в Redis
mock_save_redis.assert_awaited_once_with(123456, "ref_abc123")
async def test_capture_saves_to_redis_when_state_none(self, monkeypatch):
"""Тест сохранения payload в Redis когда FSM state недоступен."""
from app.middlewares.channel_checker import ChannelCheckerMiddleware
middleware = ChannelCheckerMiddleware()
mock_message = _create_mock_message("/start ref_xyz789", 999888)
with patch("app.middlewares.channel_checker.save_pending_payload_to_redis", new_callable=AsyncMock) as mock_save_redis:
await middleware._capture_start_payload(None, mock_message, None)
# Должен сохраняться в Redis даже если state=None
mock_save_redis.assert_awaited_once_with(999888, "ref_xyz789")
async def test_capture_ignores_message_without_payload(self, monkeypatch):
"""Тест что сообщение без payload игнорируется."""
from app.middlewares.channel_checker import ChannelCheckerMiddleware
middleware = ChannelCheckerMiddleware()
mock_state = AsyncMock()
mock_state.get_data = AsyncMock(return_value={})
mock_state.set_data = AsyncMock()
mock_message = _create_mock_message("/start", 123456) # Без payload
with patch("app.middlewares.channel_checker.save_pending_payload_to_redis", new_callable=AsyncMock) as mock_save_redis:
await middleware._capture_start_payload(mock_state, mock_message, None)
mock_state.set_data.assert_not_awaited()
mock_save_redis.assert_not_awaited()
async def test_capture_ignores_non_start_message(self, monkeypatch):
"""Тест что не-start сообщения игнорируются."""
from app.middlewares.channel_checker import ChannelCheckerMiddleware
middleware = ChannelCheckerMiddleware()
mock_state = AsyncMock()
mock_state.get_data = AsyncMock(return_value={})
mock_state.set_data = AsyncMock()
mock_message = _create_mock_message("/help something", 123456) # Не /start
with patch("app.middlewares.channel_checker.save_pending_payload_to_redis", new_callable=AsyncMock) as mock_save_redis:
await middleware._capture_start_payload(mock_state, mock_message, None)
mock_state.set_data.assert_not_awaited()
mock_save_redis.assert_not_awaited()
async def test_capture_does_not_overwrite_same_payload(self, monkeypatch):
"""Тест что одинаковый payload не перезаписывается в FSM state."""
from app.middlewares.channel_checker import ChannelCheckerMiddleware
middleware = ChannelCheckerMiddleware()
mock_state = AsyncMock()
mock_state.get_data = AsyncMock(return_value={"pending_start_payload": "ref_same"})
mock_state.set_data = AsyncMock()
mock_message = _create_mock_message("/start ref_same", 123456) # Тот же payload
with patch("app.middlewares.channel_checker.save_pending_payload_to_redis", new_callable=AsyncMock) as mock_save_redis:
await middleware._capture_start_payload(mock_state, mock_message, None)
# FSM state не должен перезаписываться
mock_state.set_data.assert_not_awaited()
# Но в Redis всё равно сохраняем (для надёжности)
mock_save_redis.assert_awaited_once()
class TestPayloadIntegration:
"""Интеграционные тесты для потока сохранения/восстановления payload."""
async def test_full_flow_fsm_state_works(self, monkeypatch):
"""Тест полного потока когда FSM state работает корректно."""
from app.middlewares.channel_checker import ChannelCheckerMiddleware
middleware = ChannelCheckerMiddleware()
# Сохраняем payload
state_storage = {}
mock_state = AsyncMock()
mock_state.get_data = AsyncMock(return_value=state_storage)
mock_state.set_data = AsyncMock(side_effect=lambda d: state_storage.update(d))
mock_message = _create_mock_message("/start ref_flow_test", 111222)
with patch("app.middlewares.channel_checker.save_pending_payload_to_redis", new_callable=AsyncMock):
await middleware._capture_start_payload(mock_state, mock_message, None)
# Проверяем что payload сохранён
assert state_storage.get("pending_start_payload") == "ref_flow_test"
async def test_payload_retrieved_from_redis_fallback(self, monkeypatch):
"""Тест что payload восстанавливается из Redis если в FSM state его нет."""
from app.middlewares.channel_checker import get_pending_payload_from_redis
mock_redis = AsyncMock()
mock_redis.get = AsyncMock(return_value=b"ref_from_redis")
mock_redis.aclose = AsyncMock()
with patch("app.middlewares.channel_checker.aioredis") as mock_aioredis:
mock_aioredis.from_url = MagicMock(return_value=mock_redis)
result = await get_pending_payload_from_redis(333444)
assert result == "ref_from_redis"