2017-10-04 22:02:45 +03:00
|
|
|
from threading import Lock
|
|
|
|
|
2017-10-05 14:01:00 +03:00
|
|
|
import re
|
|
|
|
|
2017-10-04 22:02:45 +03:00
|
|
|
from .. import utils
|
|
|
|
from ..tl import TLObject
|
2017-10-06 22:42:04 +03:00
|
|
|
from ..tl.types import (
|
|
|
|
User, Chat, Channel, PeerUser, PeerChat, PeerChannel,
|
|
|
|
InputPeerUser, InputPeerChat, InputPeerChannel
|
|
|
|
)
|
2017-10-04 22:02:45 +03:00
|
|
|
|
|
|
|
|
|
|
|
class EntityDatabase:
|
2017-10-05 14:06:51 +03:00
|
|
|
def __init__(self, input_list=None, enabled=True, enabled_full=True):
|
|
|
|
"""Creates a new entity database with an initial load of "Input"
|
|
|
|
entities, if any.
|
|
|
|
|
|
|
|
If 'enabled', input entities will be saved. The whole entity
|
|
|
|
will be saved if both 'enabled' and 'enabled_full' are True.
|
|
|
|
"""
|
2017-10-04 22:02:45 +03:00
|
|
|
self.enabled = enabled
|
2017-10-05 14:06:51 +03:00
|
|
|
self.enabled_full = enabled_full
|
2017-10-04 22:02:45 +03:00
|
|
|
|
|
|
|
self._lock = Lock()
|
|
|
|
self._entities = {} # marked_id: user|chat|channel
|
|
|
|
|
|
|
|
if input_list:
|
|
|
|
self._input_entities = {k: v for k, v in input_list}
|
|
|
|
else:
|
|
|
|
self._input_entities = {} # marked_id: hash
|
|
|
|
|
|
|
|
# TODO Allow disabling some extra mappings
|
|
|
|
self._username_id = {} # username: marked_id
|
2017-10-05 14:01:00 +03:00
|
|
|
self._phone_id = {} # phone: marked_id
|
2017-10-04 22:02:45 +03:00
|
|
|
|
|
|
|
def process(self, tlobject):
|
|
|
|
"""Processes all the found entities on the given TLObject,
|
|
|
|
unless .enabled is False.
|
|
|
|
|
|
|
|
Returns True if new input entities were added.
|
|
|
|
"""
|
|
|
|
if not self.enabled:
|
|
|
|
return False
|
|
|
|
|
|
|
|
# Save all input entities we know of
|
2017-10-05 13:29:52 +03:00
|
|
|
if not isinstance(tlobject, TLObject) and hasattr(tlobject, '__iter__'):
|
|
|
|
# This may be a list of users already for instance
|
|
|
|
return self.expand(tlobject)
|
|
|
|
|
2017-10-04 22:02:45 +03:00
|
|
|
entities = []
|
|
|
|
if hasattr(tlobject, 'chats') and hasattr(tlobject.chats, '__iter__'):
|
|
|
|
entities.extend(tlobject.chats)
|
|
|
|
if hasattr(tlobject, 'users') and hasattr(tlobject.users, '__iter__'):
|
|
|
|
entities.extend(tlobject.users)
|
|
|
|
|
|
|
|
return self.expand(entities)
|
|
|
|
|
|
|
|
def expand(self, entities):
|
|
|
|
"""Adds new input entities to the local database unconditionally.
|
|
|
|
Unknown types will be ignored.
|
|
|
|
"""
|
|
|
|
if not entities or not self.enabled:
|
|
|
|
return False
|
|
|
|
|
|
|
|
new = [] # Array of entities (User, Chat, or Channel)
|
|
|
|
new_input = {} # Dictionary of {entity_marked_id: access_hash}
|
|
|
|
for e in entities:
|
|
|
|
if not isinstance(e, TLObject):
|
|
|
|
continue
|
|
|
|
|
|
|
|
try:
|
2017-10-05 14:14:54 +03:00
|
|
|
p = utils.get_input_peer(e, allow_self=False)
|
2017-10-04 22:02:45 +03:00
|
|
|
new_input[utils.get_peer_id(p, add_mark=True)] = \
|
|
|
|
getattr(p, 'access_hash', 0) # chats won't have hash
|
|
|
|
|
2017-10-05 14:06:51 +03:00
|
|
|
if self.enabled_full:
|
|
|
|
if isinstance(e, User) \
|
|
|
|
or isinstance(e, Chat) \
|
|
|
|
or isinstance(e, Channel):
|
|
|
|
new.append(e)
|
2017-10-04 22:02:45 +03:00
|
|
|
except ValueError:
|
|
|
|
pass
|
|
|
|
|
|
|
|
with self._lock:
|
|
|
|
before = len(self._input_entities)
|
|
|
|
self._input_entities.update(new_input)
|
|
|
|
for e in new:
|
|
|
|
self._add_full_entity(e)
|
|
|
|
return len(self._input_entities) != before
|
|
|
|
|
|
|
|
def _add_full_entity(self, entity):
|
2017-10-05 14:06:51 +03:00
|
|
|
"""Adds a "full" entity (User, Chat or Channel, not "Input*"),
|
|
|
|
despite the value of self.enabled and self.enabled_full.
|
2017-10-04 22:02:45 +03:00
|
|
|
|
|
|
|
Not to be confused with UserFull, ChatFull, or ChannelFull,
|
|
|
|
"full" means simply not "Input*".
|
|
|
|
"""
|
|
|
|
marked_id = utils.get_peer_id(
|
2017-10-05 14:14:54 +03:00
|
|
|
utils.get_input_peer(entity, allow_self=False), add_mark=True
|
2017-10-04 22:02:45 +03:00
|
|
|
)
|
|
|
|
try:
|
|
|
|
old_entity = self._entities[marked_id]
|
|
|
|
old_entity.__dict__.update(entity.__dict__) # Keep old references
|
|
|
|
|
2017-10-05 14:01:00 +03:00
|
|
|
# Update must delete old username and phone
|
2017-10-04 22:02:45 +03:00
|
|
|
username = getattr(old_entity, 'username', None)
|
|
|
|
if username:
|
|
|
|
del self._username_id[username.lower()]
|
2017-10-05 14:01:00 +03:00
|
|
|
|
|
|
|
phone = getattr(old_entity, 'phone', None)
|
|
|
|
if phone:
|
|
|
|
del self._phone_id[phone]
|
2017-10-04 22:02:45 +03:00
|
|
|
except KeyError:
|
|
|
|
# Add new entity
|
|
|
|
self._entities[marked_id] = entity
|
|
|
|
|
2017-10-05 14:01:00 +03:00
|
|
|
# Always update username or phone if any
|
2017-10-04 22:02:45 +03:00
|
|
|
username = getattr(entity, 'username', None)
|
|
|
|
if username:
|
|
|
|
self._username_id[username.lower()] = marked_id
|
|
|
|
|
2017-10-05 14:01:00 +03:00
|
|
|
phone = getattr(entity, 'phone', None)
|
|
|
|
if phone:
|
|
|
|
self._username_id[phone] = marked_id
|
|
|
|
|
2017-10-04 22:02:45 +03:00
|
|
|
def __getitem__(self, key):
|
|
|
|
"""Accepts a digit only string as phone number,
|
|
|
|
otherwise it's treated as an username.
|
|
|
|
|
|
|
|
If an integer is given, it's treated as the ID of the desired User.
|
|
|
|
The ID given won't try to be guessed as the ID of a chat or channel,
|
|
|
|
as there may be an user with that ID, and it would be unreliable.
|
|
|
|
|
|
|
|
If a Peer is given (PeerUser, PeerChat, PeerChannel),
|
|
|
|
its specific entity is retrieved as User, Chat or Channel.
|
|
|
|
Note that megagroups are channels with .megagroup = True.
|
|
|
|
"""
|
|
|
|
if isinstance(key, str):
|
2017-10-05 14:01:00 +03:00
|
|
|
phone = EntityDatabase.parse_phone(key)
|
|
|
|
if phone:
|
|
|
|
return self._phone_id[phone]
|
|
|
|
else:
|
|
|
|
key = key.lstrip('@').lower()
|
|
|
|
return self._entities[self._username_id[key]]
|
2017-10-04 22:02:45 +03:00
|
|
|
|
|
|
|
if isinstance(key, int):
|
|
|
|
return self._entities[key] # normal IDs are assumed users
|
|
|
|
|
2017-10-05 13:28:04 +03:00
|
|
|
if isinstance(key, TLObject):
|
|
|
|
sc = type(key).SUBCLASS_OF_ID
|
|
|
|
if sc == 0x2d45687:
|
|
|
|
# Subclass of "Peer"
|
|
|
|
return self._entities[utils.get_peer_id(key, add_mark=True)]
|
|
|
|
elif sc in {0x2da17977, 0xc5af5d94, 0x6d44b7db}:
|
|
|
|
# Subclass of "User", "Chat" or "Channel"
|
|
|
|
return key
|
2017-10-04 22:02:45 +03:00
|
|
|
|
|
|
|
raise KeyError(key)
|
|
|
|
|
|
|
|
def __delitem__(self, key):
|
|
|
|
target = self[key]
|
|
|
|
del self._entities[key]
|
|
|
|
if getattr(target, 'username'):
|
|
|
|
del self._username_id[target.username]
|
|
|
|
|
|
|
|
# TODO Allow search by name by tokenizing the input and return a list
|
|
|
|
|
2017-10-05 14:01:00 +03:00
|
|
|
@staticmethod
|
|
|
|
def parse_phone(phone):
|
|
|
|
"""Parses the given phone, or returns None if it's invalid"""
|
|
|
|
if isinstance(phone, int):
|
|
|
|
return str(phone)
|
|
|
|
else:
|
|
|
|
phone = re.sub(r'[+()\s-]', '', str(phone))
|
|
|
|
if phone.isdigit():
|
|
|
|
return phone
|
|
|
|
|
2017-10-05 13:27:05 +03:00
|
|
|
def get_input_entity(self, peer):
|
|
|
|
try:
|
2017-10-06 22:42:04 +03:00
|
|
|
i, k = utils.get_peer_id(peer, add_mark=True, get_kind=True)
|
|
|
|
h = self._input_entities[i]
|
|
|
|
if k == PeerUser:
|
|
|
|
return InputPeerUser(i, h)
|
|
|
|
elif k == PeerChat:
|
|
|
|
return InputPeerChat(i)
|
|
|
|
elif k == PeerChannel:
|
|
|
|
return InputPeerChannel(i, h)
|
|
|
|
|
2017-10-05 13:27:05 +03:00
|
|
|
except ValueError as e:
|
|
|
|
raise KeyError(peer) from e
|
2017-10-06 22:42:04 +03:00
|
|
|
raise KeyError(peer)
|
2017-10-05 13:27:05 +03:00
|
|
|
|
2017-10-04 22:02:45 +03:00
|
|
|
def get_input_list(self):
|
|
|
|
return list(self._input_entities.items())
|
|
|
|
|
|
|
|
def clear(self, target=None):
|
|
|
|
if target is None:
|
|
|
|
self._entities.clear()
|
|
|
|
else:
|
|
|
|
del self[target]
|