fix(bot): resolve required channel join link
This commit is contained in:
@@ -9,6 +9,7 @@ from bot.middlewares.channel_subscription import ChannelSubscriptionMiddleware
|
||||
from bot.utils.channel_subscription import (
|
||||
is_required_channel_access_error,
|
||||
normalize_required_channel_id,
|
||||
normalize_required_channel_link,
|
||||
)
|
||||
|
||||
|
||||
@@ -21,10 +22,13 @@ class I18nStub:
|
||||
|
||||
|
||||
class FakeBot:
|
||||
def __init__(self, *, status="member", error=None):
|
||||
def __init__(self, *, status="member", error=None, chat=None, chat_error=None):
|
||||
self.status = status
|
||||
self.error = error
|
||||
self.chat = chat
|
||||
self.chat_error = chat_error
|
||||
self.calls = []
|
||||
self.get_chat_calls = []
|
||||
|
||||
async def get_chat_member(self, chat_id, user_id):
|
||||
self.calls.append((chat_id, user_id))
|
||||
@@ -32,11 +36,17 @@ class FakeBot:
|
||||
raise self.error
|
||||
return SimpleNamespace(status=self.status)
|
||||
|
||||
async def get_chat(self, chat_id):
|
||||
self.get_chat_calls.append(chat_id)
|
||||
if self.chat_error:
|
||||
raise self.chat_error
|
||||
return self.chat or SimpleNamespace(username="required_channel")
|
||||
|
||||
def _settings(required_channel_id):
|
||||
|
||||
def _settings(required_channel_id, required_channel_link="https://t.me/example"):
|
||||
return SimpleNamespace(
|
||||
REQUIRED_CHANNEL_ID=required_channel_id,
|
||||
REQUIRED_CHANNEL_LINK="https://t.me/example",
|
||||
REQUIRED_CHANNEL_LINK=required_channel_link,
|
||||
ADMIN_IDS=[],
|
||||
DEFAULT_LANGUAGE="en",
|
||||
)
|
||||
@@ -70,6 +80,24 @@ class RequiredChannelIdNormalizationTests(unittest.TestCase):
|
||||
self.assertIsNone(normalize_required_channel_id(""))
|
||||
self.assertIsNone(normalize_required_channel_id(0))
|
||||
|
||||
def test_normalizes_channel_links_for_join_button(self):
|
||||
self.assertEqual(
|
||||
normalize_required_channel_link("@required_channel"), "https://t.me/required_channel"
|
||||
)
|
||||
self.assertEqual(
|
||||
normalize_required_channel_link("required_channel"), "https://t.me/required_channel"
|
||||
)
|
||||
self.assertEqual(
|
||||
normalize_required_channel_link("t.me/required_channel"),
|
||||
"https://t.me/required_channel",
|
||||
)
|
||||
self.assertEqual(
|
||||
normalize_required_channel_link("https://t.me/required_channel"),
|
||||
"https://t.me/required_channel",
|
||||
)
|
||||
self.assertEqual(normalize_required_channel_link("+inviteHash"), "https://t.me/+inviteHash")
|
||||
self.assertIsNone(normalize_required_channel_link("not a valid link"))
|
||||
|
||||
def test_detects_channel_configuration_errors(self):
|
||||
self.assertTrue(
|
||||
is_required_channel_access_error(
|
||||
@@ -124,6 +152,33 @@ class RequiredChannelSubscriptionCheckTests(unittest.IsolatedAsyncioTestCase):
|
||||
event.answer.assert_awaited_once_with("check failed")
|
||||
update_user.assert_not_awaited()
|
||||
|
||||
async def test_join_button_prefers_channel_link_resolved_from_required_id(self):
|
||||
bot = FakeBot(
|
||||
status="left",
|
||||
chat=SimpleNamespace(username="required_channel", invite_link=None),
|
||||
)
|
||||
event = _message_event(bot)
|
||||
user = _db_user()
|
||||
|
||||
with patch("bot.handlers.user.start.user_dal.update_user", AsyncMock()):
|
||||
result = await ensure_required_channel_subscription(
|
||||
event,
|
||||
_settings(1234567890, required_channel_link="https://t.me/main_sales_bot"),
|
||||
I18nStub(),
|
||||
"en",
|
||||
AsyncMock(),
|
||||
db_user=user,
|
||||
)
|
||||
|
||||
self.assertFalse(result)
|
||||
self.assertEqual(bot.calls, [(-1001234567890, 42)])
|
||||
self.assertEqual(bot.get_chat_calls, [-1001234567890])
|
||||
reply_markup = event.answer.await_args.kwargs["reply_markup"]
|
||||
self.assertEqual(
|
||||
reply_markup.inline_keyboard[0][0].url,
|
||||
"https://t.me/required_channel",
|
||||
)
|
||||
|
||||
|
||||
class ChannelSubscriptionMiddlewareTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_middleware_accepts_cached_verification_for_normalized_channel_id(self):
|
||||
@@ -146,6 +201,40 @@ class ChannelSubscriptionMiddlewareTests(unittest.IsolatedAsyncioTestCase):
|
||||
self.assertEqual(result, "ok")
|
||||
handler.assert_awaited_once_with(event, data)
|
||||
|
||||
async def test_middleware_prompt_uses_channel_link_resolved_from_required_id(self):
|
||||
middleware = ChannelSubscriptionMiddleware(
|
||||
_settings(1234567890, required_channel_link="https://t.me/main_sales_bot"),
|
||||
I18nStub(),
|
||||
)
|
||||
handler = AsyncMock(return_value="ok")
|
||||
message = SimpleNamespace(text="menu", answer=AsyncMock())
|
||||
event = SimpleNamespace(callback_query=None, message=message)
|
||||
bot = FakeBot(
|
||||
chat=SimpleNamespace(username="required_channel", invite_link=None),
|
||||
)
|
||||
data = {
|
||||
"bot": bot,
|
||||
"event_from_user": SimpleNamespace(id=42),
|
||||
"session": AsyncMock(),
|
||||
"i18n_data": {"current_language": "en", "i18n_instance": I18nStub()},
|
||||
}
|
||||
user = _db_user(verified=False, verified_for=None)
|
||||
|
||||
with patch(
|
||||
"bot.middlewares.channel_subscription.user_dal.get_user_by_id",
|
||||
AsyncMock(return_value=user),
|
||||
):
|
||||
result = await middleware(handler, event, data)
|
||||
|
||||
self.assertIsNone(result)
|
||||
handler.assert_not_awaited()
|
||||
self.assertEqual(bot.get_chat_calls, [-1001234567890])
|
||||
reply_markup = message.answer.await_args.kwargs["reply_markup"]
|
||||
self.assertEqual(
|
||||
reply_markup.inline_keyboard[0][0].url,
|
||||
"https://t.me/required_channel",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user