Add more type hints

This commit is contained in:
Tulir Asokan
2018-07-25 10:40:31 -04:00
parent ae334b9a04
commit dbfb980bde
20 changed files with 751 additions and 595 deletions
+81 -44
View File
@@ -14,26 +14,48 @@
#
# You should have received a copy of the GNU Affero General Public License
# along with this program. If not, see <https://www.gnu.org/licenses/>.
from typing import Tuple, Optional, List, Union, TYPE_CHECKING
from abc import ABC, abstractmethod
import asyncio
import logging
import platform
from telethon.tl.types import *
from mautrix_appservice import MatrixRequestError
from sqlalchemy import orm
from telethon.tl.types import Channel, ChannelForbidden, Chat, ChatForbidden, Message, \
MessageActionChannelMigrateFrom, MessageService, PeerUser, TypeUpdate, \
UpdateChannelPinnedMessage, UpdateChatAdmins, UpdateChatParticipantAdmin, \
UpdateChatParticipants, UpdateChatUserTyping, UpdateDeleteChannelMessages, \
UpdateDeleteMessages, UpdateEditChannelMessage, UpdateEditMessage, UpdateNewChannelMessage, \
UpdateNewMessage, UpdateReadHistoryOutbox, UpdateShortChatMessage, UpdateShortMessage, \
UpdateUserName, UpdateUserPhoto, UpdateUserStatus, UpdateUserTyping, User, UserStatusOffline, \
UserStatusOnline
from mautrix_appservice import MatrixRequestError, AppService
from alchemysession import AlchemySessionContainer
from .tgclient import MautrixTelegramClient
from .db import Message as DBMessage
from . import portal as po, puppet as pu, __version__
from .db import Message as DBMessage
from .tgclient import MautrixTelegramClient
config = None
if TYPE_CHECKING:
from .context import Context
from .config import Config
config = None # type: Config
# Value updated from config in init()
MAX_DELETIONS = 10
MAX_DELETIONS = 10 # type: int
UpdateMessage = Union[UpdateShortChatMessage, UpdateShortMessage, UpdateNewChannelMessage,
UpdateNewMessage, UpdateEditMessage, UpdateEditChannelMessage]
UpdateMessageContent = Union[UpdateShortMessage, UpdateShortChatMessage, Message, MessageService]
class AbstractUser:
session_container = None
loop = None
log = None
db = None
az = None
class AbstractUser(ABC):
session_container = None # type: AlchemySessionContainer
loop = None # type: asyncio.AbstractEventLoop
log = None # type: logging.Logger
db = None # type: orm.Session
az = None # type: AppService
def __init__(self):
self.puppet_whitelisted = False # type: bool
@@ -47,22 +69,22 @@ class AbstractUser:
self.is_bot = False # type: bool
@property
def connected(self):
def connected(self) -> bool:
return self.client and self.client.is_connected()
@property
def _proxy_settings(self):
type = config["telegram.proxy.type"].lower()
if type == "disabled":
def _proxy_settings(self) -> Optional[Tuple[int, str, str, str, str, str]]:
proxy_type = config["telegram.proxy.type"].lower()
if proxy_type == "disabled":
return None
elif type == "socks4":
type = 1
elif type == "socks5":
type = 2
elif type == "http":
type = 3
elif proxy_type == "socks4":
proxy_type = 1
elif proxy_type == "socks5":
proxy_type = 2
elif proxy_type == "http":
proxy_type = 3
return (type,
return (proxy_type,
config["telegram.proxy.address"], config["telegram.proxy.port"],
config["telegram.proxy.rdns"],
config["telegram.proxy.username"], config["telegram.proxy.password"])
@@ -83,20 +105,30 @@ class AbstractUser:
proxy=self._proxy_settings)
self.client.add_event_handler(self._update_catch)
async def update(self, update):
@abstractmethod
async def update(self, update: TypeUpdate) -> bool:
return False
@abstractmethod
async def post_login(self):
raise NotImplementedError()
async def _update_catch(self, update):
@abstractmethod
def register_portal(self, portal: po.Portal):
raise NotImplementedError()
@abstractmethod
def unregister_portal(self, portal: po.Portal):
raise NotImplementedError()
async def _update_catch(self, update: TypeUpdate):
try:
if not await self.update(update):
await self._update(update)
except Exception:
self.log.exception("Failed to handle Telegram update")
async def get_dialogs(self, limit=None) -> List[Union[Chat, Channel]]:
async def get_dialogs(self, limit: int = None) -> List[Union[Chat, Channel]]:
if self.is_bot:
return []
dialogs = await self.client.get_dialogs(limit=limit)
@@ -106,18 +138,19 @@ class AbstractUser:
and (dialog.entity.deactivated or dialog.entity.left)))]
@property
def name(self):
@abstractmethod
def name(self) -> str:
raise NotImplementedError()
async def is_logged_in(self):
async def is_logged_in(self) -> bool:
return self.client and await self.client.is_user_authorized()
async def has_full_access(self, allow_bot=False):
async def has_full_access(self, allow_bot: bool = False) -> bool:
return (self.puppet_whitelisted
and (not self.is_bot or allow_bot)
and await self.is_logged_in())
async def start(self, delete_unless_authenticated=False):
async def start(self, delete_unless_authenticated: bool = False) -> "AbstractUser":
if not self.client:
self._init_client()
await self.client.connect()
@@ -144,7 +177,7 @@ class AbstractUser:
# region Telegram update handling
async def _update(self, update):
async def _update(self, update: TypeUpdate):
if isinstance(update, (UpdateShortChatMessage, UpdateShortMessage, UpdateNewChannelMessage,
UpdateNewMessage, UpdateEditMessage, UpdateEditChannelMessage)):
await self.update_message(update)
@@ -169,17 +202,19 @@ class AbstractUser:
else:
self.log.debug("Unhandled update: %s", update)
async def update_pinned_messages(self, update):
@staticmethod
async def update_pinned_messages(update: UpdateChannelPinnedMessage):
portal = po.Portal.get_by_tgid(update.channel_id)
if portal and portal.mxid:
await portal.receive_telegram_pin_id(update.id)
async def update_participants(self, update):
@staticmethod
async def update_participants(update: UpdateChatParticipants):
portal = po.Portal.get_by_tgid(update.participants.chat_id)
if portal and portal.mxid:
await portal.update_telegram_participants(update.participants.participants)
async def update_read_receipt(self, update):
async def update_read_receipt(self, update: UpdateReadHistoryOutbox):
if not isinstance(update.peer, PeerUser):
self.log.debug("Unexpected read receipt peer: %s", update.peer)
return
@@ -196,7 +231,7 @@ class AbstractUser:
puppet = pu.Puppet.get(update.peer.user_id)
await puppet.intent.mark_read(portal.mxid, message.mxid)
async def update_admin(self, update):
async def update_admin(self, update: Union[UpdateChatAdmins, UpdateChatParticipantAdmin]):
# TODO duplication not checked
portal = po.Portal.get_by_tgid(update.chat_id, peer_type="chat")
if isinstance(update, UpdateChatAdmins):
@@ -206,7 +241,7 @@ class AbstractUser:
else:
self.log.warning("Unexpected admin status update: %s", update)
async def update_typing(self, update):
async def update_typing(self, update: Union[UpdateUserTyping, UpdateChatUserTyping]):
if isinstance(update, UpdateUserTyping):
portal = po.Portal.get_by_tgid(update.user_id, self.tgid, "user")
else:
@@ -214,7 +249,7 @@ class AbstractUser:
sender = pu.Puppet.get(update.user_id)
await portal.handle_telegram_typing(sender, update)
async def update_others_info(self, update):
async def update_others_info(self, update: Union[UpdateUserName, UpdateUserPhoto]):
# TODO duplication not checked
puppet = pu.Puppet.get(update.user_id)
if isinstance(update, UpdateUserName):
@@ -226,7 +261,7 @@ class AbstractUser:
else:
self.log.warning("Unexpected other user info update: %s", update)
async def update_status(self, update):
async def update_status(self, update: UpdateUserStatus):
puppet = pu.Puppet.get(update.user_id)
if isinstance(update.status, UserStatusOnline):
await puppet.default_mxid_intent.set_presence("online")
@@ -236,7 +271,9 @@ class AbstractUser:
self.log.warning("Unexpected user status update: %s", update)
return
def get_message_details(self, update):
def get_message_details(self, update: UpdateMessage) -> Tuple[UpdateMessageContent,
Optional[pu.Puppet],
Optional[po.Portal]]:
if isinstance(update, UpdateShortChatMessage):
portal = po.Portal.get_by_tgid(update.chat_id, peer_type="chat")
sender = pu.Puppet.get(update.from_id)
@@ -259,7 +296,7 @@ class AbstractUser:
return update, sender, portal
@staticmethod
async def _try_redact(portal, message):
async def _try_redact(portal: po.Portal, message: DBMessage):
if not portal:
return
try:
@@ -267,7 +304,7 @@ class AbstractUser:
except MatrixRequestError:
pass
async def delete_message(self, update):
async def delete_message(self, update: UpdateDeleteMessages):
if len(update.messages) > MAX_DELETIONS:
return
@@ -283,7 +320,7 @@ class AbstractUser:
await self._try_redact(portal, message)
self.db.commit()
async def delete_channel_message(self, update):
async def delete_channel_message(self, update: UpdateDeleteChannelMessages):
if len(update.messages) > MAX_DELETIONS:
return
@@ -299,7 +336,7 @@ class AbstractUser:
await self._try_redact(portal, message)
self.db.commit()
async def update_message(self, original_update):
async def update_message(self, original_update: UpdateMessage):
update, sender, portal = self.get_message_details(original_update)
if isinstance(update, MessageService):
@@ -325,7 +362,7 @@ class AbstractUser:
# endregion
def init(context):
def init(context: "Context"):
global config, MAX_DELETIONS
AbstractUser.az, AbstractUser.db, config, AbstractUser.loop, _ = context
AbstractUser.session_container = context.session_container