fix: normalize required channel checks

This commit is contained in:
3252a8
2026-06-01 22:13:35 +03:00
parent 5b19ba2c2f
commit 5218ede0f1
5 changed files with 222 additions and 2 deletions
+28 -1
View File
@@ -24,6 +24,10 @@ from bot.services.referral_service import ReferralService
from bot.services.subscription_service import SubscriptionService from bot.services.subscription_service import SubscriptionService
from bot.services.telegram_notifications import TELEGRAM_NOTIFICATIONS_ENABLED from bot.services.telegram_notifications import TELEGRAM_NOTIFICATIONS_ENABLED
from bot.utils.callback_answer import safe_answer_callback from bot.utils.callback_answer import safe_answer_callback
from bot.utils.channel_subscription import (
is_required_channel_access_error,
normalize_required_channel_id,
)
from bot.utils.install_links import ( from bot.utils.install_links import (
append_install_share_link_text, append_install_share_link_text,
ensure_user_install_guide_links, ensure_user_install_guide_links,
@@ -215,7 +219,7 @@ async def ensure_required_channel_subscription(
Verify that the user is a member of the required channel (if configured). Verify that the user is a member of the required channel (if configured).
Returns True when access can proceed, False when user must subscribe first. Returns True when access can proceed, False when user must subscribe first.
""" """
required_channel_id = settings.REQUIRED_CHANNEL_ID required_channel_id = normalize_required_channel_id(settings.REQUIRED_CHANNEL_ID)
if not required_channel_id: if not required_channel_id:
return True return True
@@ -279,6 +283,29 @@ async def ensure_required_channel_subscription(
if status_value in allowed_statuses: if status_value in allowed_statuses:
is_member = True is_member = True
except TelegramBadRequest as bad_request: except TelegramBadRequest as bad_request:
if is_required_channel_access_error(bad_request):
logging.error(
"Required channel check failed due to channel access/configuration error "
"(configured=%s, normalized=%s): %s",
settings.REQUIRED_CHANNEL_ID,
required_channel_id,
bad_request,
)
error_text = translate("channel_subscription_check_failed")
if isinstance(event, types.CallbackQuery):
try:
await event.answer(error_text, show_alert=True)
except Exception:
pass
if message_obj:
try:
await message_obj.answer(error_text)
except Exception:
pass
else:
await event.answer(error_text)
return False
logging.info( logging.info(
"Required channel check: user %s not subscribed (details: %s)", "Required channel check: user %s not subscribed (details: %s)",
user_id, user_id,
@@ -11,6 +11,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from bot.keyboards.inline.user_keyboards import get_channel_subscription_keyboard from bot.keyboards.inline.user_keyboards import get_channel_subscription_keyboard
from bot.middlewares.i18n import JsonI18n from bot.middlewares.i18n import JsonI18n
from bot.utils.channel_subscription import normalize_required_channel_id
from config.settings import Settings from config.settings import Settings
from db.dal import user_dal from db.dal import user_dal
@@ -32,7 +33,7 @@ class ChannelSubscriptionMiddleware(BaseMiddleware):
event: Update, event: Update,
data: Dict[str, Any], data: Dict[str, Any],
) -> Any: ) -> Any:
required_channel_id = self.settings.REQUIRED_CHANNEL_ID required_channel_id = normalize_required_channel_id(self.settings.REQUIRED_CHANNEL_ID)
if not required_channel_id: if not required_channel_id:
return await handler(event, data) return await handler(event, data)
+40
View File
@@ -0,0 +1,40 @@
from typing import Optional
def normalize_required_channel_id(value: object) -> Optional[int]:
if value is None:
return None
raw = str(value).strip()
if not raw:
return None
try:
channel_id = int(raw)
except (TypeError, ValueError):
return None
if channel_id == 0:
return None
if channel_id > 0:
return int(f"-100{channel_id}")
raw_abs = str(abs(channel_id))
if raw.startswith("-100"):
return channel_id
if abs(channel_id) < 1_000_000_000:
return channel_id
return -int(f"100{raw_abs}")
def is_required_channel_access_error(error: BaseException) -> bool:
message = str(error).lower()
configuration_markers = (
"chat not found",
"bot is not a member",
"not enough rights",
"have no rights",
"kicked",
)
return any(marker in message for marker in configuration_markers)
+1
View File
@@ -1018,6 +1018,7 @@ class Settings(BaseSettings):
"LOG_SUPPORT_THREAD_ID", "LOG_SUPPORT_THREAD_ID",
"BACKUP_CHAT_ID", "BACKUP_CHAT_ID",
"BACKUP_THREAD_ID", "BACKUP_THREAD_ID",
"REQUIRED_CHANNEL_ID",
mode="before", mode="before",
) )
@classmethod @classmethod
+151
View File
@@ -0,0 +1,151 @@
import unittest
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
from aiogram.exceptions import TelegramBadRequest
from bot.handlers.user.start import ensure_required_channel_subscription
from bot.middlewares.channel_subscription import ChannelSubscriptionMiddleware
from bot.utils.channel_subscription import (
is_required_channel_access_error,
normalize_required_channel_id,
)
class I18nStub:
def gettext(self, _lang, key, **_kwargs):
return {
"channel_subscription_required": "subscribe first",
"channel_subscription_check_failed": "check failed",
}.get(key, key)
class FakeBot:
def __init__(self, *, status="member", error=None):
self.status = status
self.error = error
self.calls = []
async def get_chat_member(self, chat_id, user_id):
self.calls.append((chat_id, user_id))
if self.error:
raise self.error
return SimpleNamespace(status=self.status)
def _settings(required_channel_id):
return SimpleNamespace(
REQUIRED_CHANNEL_ID=required_channel_id,
REQUIRED_CHANNEL_LINK="https://t.me/example",
ADMIN_IDS=[],
DEFAULT_LANGUAGE="en",
)
def _message_event(bot, user_id=42):
return SimpleNamespace(
from_user=SimpleNamespace(id=user_id),
bot=bot,
answer=AsyncMock(),
)
def _db_user(*, verified=False, verified_for=None):
return SimpleNamespace(
channel_subscription_verified=verified,
channel_subscription_verified_for=verified_for,
)
class RequiredChannelIdNormalizationTests(unittest.TestCase):
def test_normalizes_raw_channel_ids_to_bot_api_chat_ids(self):
self.assertEqual(normalize_required_channel_id(1234567890), -1001234567890)
self.assertEqual(normalize_required_channel_id("1234567890"), -1001234567890)
self.assertEqual(normalize_required_channel_id(-1234567890), -1001234567890)
self.assertEqual(normalize_required_channel_id(-1001234567890), -1001234567890)
self.assertEqual(normalize_required_channel_id(-123456789), -123456789)
def test_ignores_empty_channel_ids(self):
self.assertIsNone(normalize_required_channel_id(None))
self.assertIsNone(normalize_required_channel_id(""))
self.assertIsNone(normalize_required_channel_id(0))
def test_detects_channel_configuration_errors(self):
self.assertTrue(
is_required_channel_access_error(
TelegramBadRequest(method=None, message="Bad Request: chat not found")
)
)
self.assertFalse(
is_required_channel_access_error(
TelegramBadRequest(method=None, message="Bad Request: user not found")
)
)
class RequiredChannelSubscriptionCheckTests(unittest.IsolatedAsyncioTestCase):
async def test_check_uses_normalized_channel_id_and_persists_it(self):
bot = FakeBot(status="member")
event = _message_event(bot)
user = _db_user()
with patch("bot.handlers.user.start.user_dal.update_user", AsyncMock()) as update_user:
result = await ensure_required_channel_subscription(
event,
_settings(1234567890),
I18nStub(),
"en",
AsyncMock(),
db_user=user,
)
self.assertTrue(result)
self.assertEqual(bot.calls, [(-1001234567890, 42)])
payload = update_user.await_args.args[2]
self.assertTrue(payload["channel_subscription_verified"])
self.assertEqual(payload["channel_subscription_verified_for"], -1001234567890)
async def test_channel_access_errors_show_check_failed_without_persisting_false_status(self):
bot = FakeBot(error=TelegramBadRequest(method=None, message="Bad Request: chat not found"))
event = _message_event(bot)
user = _db_user()
with patch("bot.handlers.user.start.user_dal.update_user", AsyncMock()) as update_user:
result = await ensure_required_channel_subscription(
event,
_settings(1234567890),
I18nStub(),
"en",
AsyncMock(),
db_user=user,
)
self.assertFalse(result)
event.answer.assert_awaited_once_with("check failed")
update_user.assert_not_awaited()
class ChannelSubscriptionMiddlewareTests(unittest.IsolatedAsyncioTestCase):
async def test_middleware_accepts_cached_verification_for_normalized_channel_id(self):
middleware = ChannelSubscriptionMiddleware(_settings(1234567890), I18nStub())
handler = AsyncMock(return_value="ok")
event = SimpleNamespace(callback_query=None, message=None)
data = {
"event_from_user": SimpleNamespace(id=42),
"session": AsyncMock(),
"i18n_data": {"current_language": "en", "i18n_instance": I18nStub()},
}
user = _db_user(verified=True, verified_for=-1001234567890)
with patch(
"bot.middlewares.channel_subscription.user_dal.get_user_by_id",
AsyncMock(return_value=user),
):
result = await middleware(handler, event, data)
self.assertEqual(result, "ok")
handler.assert_awaited_once_with(event, data)
if __name__ == "__main__":
unittest.main()