diff --git a/pyrogram/client.py b/pyrogram/client.py index b88ea6918c6..36e165929a4 100644 --- a/pyrogram/client.py +++ b/pyrogram/client.py @@ -60,6 +60,20 @@ log = logging.getLogger(__name__) +def getattr_nested(target, *args: str, default=None): + """ + TODO: move into utils? + """ + result = target + for attr in args: + if hasattr(result, attr): + result = getattr(result, attr) + else: + return default + + return result + + class Client(Methods): """Pyrogram Client, the main means for interacting with Telegram. @@ -483,79 +497,89 @@ async def fetch_peers(self, peers: List[Union[raw.types.User, raw.types.Chat, ra return is_min - async def handle_updates(self, updates): - if isinstance(updates, (raw.types.Updates, raw.types.UpdatesCombined)): - is_min = any(( - await self.fetch_peers(updates.users), - await self.fetch_peers(updates.chats), - )) + def _handle_updates__extract_channel_id(self, update): + """ + Re-implementation extraction `channel_id` logic as is + TODO: maybe it would be easier to check update type? + """ - users = {u.id: u for u in updates.users} - chats = {c.id: c for c in updates.chats} + return ( + getattr_nested(update, "message", "peer_id", "channel_id") or + getattr(update, "channel_id", None) + ) - for update in updates.updates: - channel_id = getattr( - getattr( - getattr( - update, "message", None - ), "peer_id", None - ), "channel_id", None - ) or getattr(update, "channel_id", None) + async def _handle_updates__updates(self, updates): + is_min = any(( + await self.fetch_peers(updates.users), + await self.fetch_peers(updates.chats), + )) - pts = getattr(update, "pts", None) - pts_count = getattr(update, "pts_count", None) + users = {u.id: u for u in updates.users} + chats = {c.id: c for c in updates.chats} - if isinstance(update, raw.types.UpdateChannelTooLong): - log.warning(update) + for update in updates.updates: + channel_id = self._handle_updates__extract_channel_id(update) + pts = getattr(update, "pts", None) + pts_count = getattr(update, "pts_count", None) - if isinstance(update, raw.types.UpdateNewChannelMessage) and is_min: - message = update.message + if isinstance(update, raw.types.UpdateChannelTooLong): + log.warning(update) - if not isinstance(message, raw.types.MessageEmpty): - try: - diff = await self.invoke( - raw.functions.updates.GetChannelDifference( - channel=await self.resolve_peer(utils.get_channel_id(channel_id)), - filter=raw.types.ChannelMessagesFilter( - ranges=[raw.types.MessageRange( - min_id=update.message.id, - max_id=update.message.id - )] - ), - pts=pts - pts_count, - limit=pts - ) - ) - except ChannelPrivate: - pass - else: - if not isinstance(diff, raw.types.updates.ChannelDifferenceEmpty): - users.update({u.id: u for u in diff.users}) - chats.update({c.id: c for c in diff.chats}) + if isinstance(update, raw.types.UpdateNewChannelMessage) and is_min: + message = update.message - self.dispatcher.updates_queue.put_nowait((update, users, chats)) - elif isinstance(updates, (raw.types.UpdateShortMessage, raw.types.UpdateShortChatMessage)): - diff = await self.invoke( - raw.functions.updates.GetDifference( - pts=updates.pts - updates.pts_count, - date=updates.date, - qts=-1 - ) + if not isinstance(message, raw.types.MessageEmpty): + try: + diff = await self.invoke( + raw.functions.updates.GetChannelDifference( + channel=await self.resolve_peer(utils.get_channel_id(channel_id)), + filter=raw.types.ChannelMessagesFilter( + ranges=[raw.types.MessageRange( + min_id=update.message.id, + max_id=update.message.id + )] + ), + pts=pts - pts_count, + limit=pts + ) + ) + except ChannelPrivate: + pass + else: + if not isinstance(diff, raw.types.updates.ChannelDifferenceEmpty): + users.update({u.id: u for u in diff.users}) + chats.update({c.id: c for c in diff.chats}) + + self.dispatcher.updates_queue.put_nowait((update, users, chats)) + + async def _handle_updates__message(self, updates): + diff = await self.invoke( + raw.functions.updates.GetDifference( + pts=updates.pts - updates.pts_count, + date=updates.date, + qts=-1 ) + ) + + if diff.new_messages: + self.dispatcher.updates_queue.put_nowait(( + raw.types.UpdateNewMessage( + message=diff.new_messages[0], + pts=updates.pts, + pts_count=updates.pts_count + ), + {u.id: u for u in diff.users}, + {c.id: c for c in diff.chats} + )) + else: + if diff.other_updates: # The other_updates list can be empty + self.dispatcher.updates_queue.put_nowait((diff.other_updates[0], {}, {})) - if diff.new_messages: - self.dispatcher.updates_queue.put_nowait(( - raw.types.UpdateNewMessage( - message=diff.new_messages[0], - pts=updates.pts, - pts_count=updates.pts_count - ), - {u.id: u for u in diff.users}, - {c.id: c for c in diff.chats} - )) - else: - if diff.other_updates: # The other_updates list can be empty - self.dispatcher.updates_queue.put_nowait((diff.other_updates[0], {}, {})) + async def handle_updates(self, updates): + if isinstance(updates, (raw.types.Updates, raw.types.UpdatesCombined)): + self._handle_updates__updates(updates) + elif isinstance(updates, (raw.types.UpdateShortMessage, raw.types.UpdateShortChatMessage)): + self._handle_updates__message(updates) elif isinstance(updates, raw.types.UpdateShort): self.dispatcher.updates_queue.put_nowait((updates.update, {}, {})) elif isinstance(updates, raw.types.UpdatesTooLong):