Add type annotations (PEP 484) (#25)
This commit is contained in:
+35
-26
@@ -4,13 +4,15 @@ import logging
|
||||
import time
|
||||
import weakref
|
||||
from enum import Enum
|
||||
from json import JSONEncoder
|
||||
from threading import Thread
|
||||
from typing import Deque, Dict, Iterable, Optional, Tuple, Type, Union
|
||||
|
||||
from .credentials import CertificateCredentials, Credentials
|
||||
from .errors import ConnectionFailed, exception_class_for_reason
|
||||
|
||||
# We don't generally need to know about the Credentials subclasses except to
|
||||
# keep the old API, where APNsClient took a cert_file
|
||||
from .credentials import CertificateCredentials
|
||||
from .payload import Payload
|
||||
|
||||
|
||||
class NotificationPriority(Enum):
|
||||
@@ -35,10 +37,14 @@ class APNsClient(object):
|
||||
DEFAULT_PORT = 443
|
||||
ALTERNATIVE_PORT = 2197
|
||||
|
||||
def __init__(self, credentials, use_sandbox=False, use_alternative_port=False, proto=None, json_encoder=None,
|
||||
password=None, proxy_host=None, proxy_port=None, heartbeat_period=None):
|
||||
if credentials is None or isinstance(credentials, str):
|
||||
self.__credentials = CertificateCredentials(credentials, password)
|
||||
def __init__(self,
|
||||
credentials: Union[Credentials, str],
|
||||
use_sandbox: bool = False, use_alternative_port: bool = False, proto: Optional[str] = None,
|
||||
json_encoder: Optional[Type[JSONEncoder]] = None, password: Optional[str] = None,
|
||||
proxy_host: Optional[str] = None, proxy_port: Optional[int] = None,
|
||||
heartbeat_period: Optional[float] = None) -> None:
|
||||
if isinstance(credentials, str):
|
||||
self.__credentials = CertificateCredentials(credentials, password) # type: Credentials
|
||||
else:
|
||||
self.__credentials = credentials
|
||||
self._init_connection(use_sandbox, use_alternative_port, proto, proxy_host, proxy_port)
|
||||
@@ -47,18 +53,19 @@ class APNsClient(object):
|
||||
self._start_heartbeat(heartbeat_period)
|
||||
|
||||
self.__json_encoder = json_encoder
|
||||
self.__max_concurrent_streams = None
|
||||
self.__max_concurrent_streams = 0
|
||||
self.__previous_server_max_concurrent_streams = None
|
||||
|
||||
def _init_connection(self, use_sandbox, use_alternative_port, proto, proxy_host, proxy_port):
|
||||
def _init_connection(self, use_sandbox: bool, use_alternative_port: bool, proto: Optional[str],
|
||||
proxy_host: Optional[str], proxy_port: Optional[int]) -> None:
|
||||
server = self.SANDBOX_SERVER if use_sandbox else self.LIVE_SERVER
|
||||
port = self.ALTERNATIVE_PORT if use_alternative_port else self.DEFAULT_PORT
|
||||
self._connection = self.__credentials.create_connection(server, port, proto, proxy_host, proxy_port)
|
||||
|
||||
def _start_heartbeat(self, heartbeat_period):
|
||||
def _start_heartbeat(self, heartbeat_period: float) -> None:
|
||||
conn_ref = weakref.ref(self._connection)
|
||||
|
||||
def watchdog():
|
||||
def watchdog() -> None:
|
||||
while True:
|
||||
conn = conn_ref()
|
||||
if conn is None:
|
||||
@@ -71,8 +78,9 @@ class APNsClient(object):
|
||||
thread.setDaemon(True)
|
||||
thread.start()
|
||||
|
||||
def send_notification(self, token_hex, notification, topic=None, priority=NotificationPriority.Immediate,
|
||||
expiration=None, collapse_id=None):
|
||||
def send_notification(self, token_hex: str, notification: Payload, topic: Optional[str] = None,
|
||||
priority: NotificationPriority = NotificationPriority.Immediate,
|
||||
expiration: Optional[int] = None, collapse_id: Optional[str] = None) -> None:
|
||||
stream_id = self.send_notification_async(token_hex, notification, topic, priority, expiration, collapse_id)
|
||||
result = self.get_notification_result(stream_id)
|
||||
if result != 'Success':
|
||||
@@ -82,8 +90,9 @@ class APNsClient(object):
|
||||
else:
|
||||
raise exception_class_for_reason(result)
|
||||
|
||||
def send_notification_async(self, token_hex, notification, topic=None, priority=NotificationPriority.Immediate,
|
||||
expiration=None, collapse_id=None):
|
||||
def send_notification_async(self, token_hex: str, notification: Payload, topic: Optional[str] = None,
|
||||
priority: NotificationPriority = NotificationPriority.Immediate,
|
||||
expiration: Optional[int] = None, collapse_id: Optional[str] = None) -> int:
|
||||
json_str = json.dumps(notification.dict(), cls=self.__json_encoder, ensure_ascii=False, separators=(',', ':'))
|
||||
json_payload = json_str.encode('utf-8')
|
||||
|
||||
@@ -105,10 +114,10 @@ class APNsClient(object):
|
||||
headers['apns-collapse-id'] = collapse_id
|
||||
|
||||
url = '/3/device/{}'.format(token_hex)
|
||||
stream_id = self._connection.request('POST', url, json_payload, headers)
|
||||
stream_id = self._connection.request('POST', url, json_payload, headers) # type: int
|
||||
return stream_id
|
||||
|
||||
def get_notification_result(self, stream_id):
|
||||
def get_notification_result(self, stream_id: int) -> Union[str, Tuple[str, str]]:
|
||||
"""
|
||||
Get result for specified stream
|
||||
The function returns: 'Success' or 'failure reason' or ('Unregistered', timestamp)
|
||||
@@ -118,14 +127,16 @@ class APNsClient(object):
|
||||
return 'Success'
|
||||
else:
|
||||
raw_data = response.read().decode('utf-8')
|
||||
data = json.loads(raw_data)
|
||||
data = json.loads(raw_data) # type: Dict[str, str]
|
||||
if response.status == 410:
|
||||
return data['reason'], data['timestamp']
|
||||
else:
|
||||
return data['reason']
|
||||
|
||||
def send_notification_batch(self, notifications, topic=None, priority=NotificationPriority.Immediate,
|
||||
expiration=None, collapse_id=None):
|
||||
def send_notification_batch(self, notifications: Iterable[Notification], topic: Optional[str] = None,
|
||||
priority: NotificationPriority = NotificationPriority.Immediate,
|
||||
expiration: Optional[int] = None, collapse_id: Optional[str] = None
|
||||
) -> Dict[str, Union[str, Tuple[str, str]]]:
|
||||
"""
|
||||
Send a notification to a list of tokens in batch. Instead of sending a synchronous request
|
||||
for each token, send multiple requests concurrently. This is done on the same connection,
|
||||
@@ -146,7 +157,7 @@ class APNsClient(object):
|
||||
self.connect()
|
||||
|
||||
results = {}
|
||||
open_streams = collections.deque()
|
||||
open_streams = collections.deque() # type: Deque[RequestStream]
|
||||
# Loop on the tokens, sending as many requests as possible concurrently to APNs.
|
||||
# When reaching the maximum concurrent streams limit, wait for a response before sending
|
||||
# another request.
|
||||
@@ -154,7 +165,7 @@ class APNsClient(object):
|
||||
# Update the max_concurrent_streams on every iteration since a SETTINGS frame can be
|
||||
# sent by the server at any time.
|
||||
self.update_max_concurrent_streams()
|
||||
if self.should_send_notification(next_notification, open_streams):
|
||||
if next_notification is not None and len(open_streams) < self.__max_concurrent_streams:
|
||||
logger.info('Sending to token %s', next_notification.token)
|
||||
stream_id = self.send_notification_async(next_notification.token, next_notification.payload, topic,
|
||||
priority, expiration, collapse_id)
|
||||
@@ -175,10 +186,7 @@ class APNsClient(object):
|
||||
|
||||
return results
|
||||
|
||||
def should_send_notification(self, notification, open_streams):
|
||||
return notification is not None and len(open_streams) < self.__max_concurrent_streams
|
||||
|
||||
def update_max_concurrent_streams(self):
|
||||
def update_max_concurrent_streams(self) -> None:
|
||||
# Get the max_concurrent_streams setting returned by the server.
|
||||
# The max_concurrent_streams value is saved in the H2Connection instance that must be
|
||||
# accessed using a with statement in order to acquire a lock.
|
||||
@@ -204,13 +212,14 @@ class APNsClient(object):
|
||||
logger.info('APNs set max_concurrent_streams to %s', max_concurrent_streams)
|
||||
self.__max_concurrent_streams = max_concurrent_streams
|
||||
|
||||
def connect(self):
|
||||
def connect(self) -> None:
|
||||
"""
|
||||
Establish a connection to APNs. If already connected, the function does nothing. If the
|
||||
connection fails, the function retries up to MAX_CONNECTION_RETRIES times.
|
||||
"""
|
||||
retries = 0
|
||||
while retries < MAX_CONNECTION_RETRIES:
|
||||
# noinspection PyBroadException
|
||||
try:
|
||||
self._connection.connect()
|
||||
logger.info('Connected to APNs')
|
||||
|
||||
Reference in New Issue
Block a user