Merge pull request #2390 from Gy9vin/main
fix(referral): исправить потерю реферальных кодов при обязательной по…
This commit is contained in:
+47
-19
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user