Support sending notifications in batch. (#27)

This commit is contained in:
Matan Gover
2017-01-23 20:10:35 +03:00
committed by Sergey Petrov
parent 1262d6c2bd
commit 701b587c32
9 changed files with 279 additions and 23 deletions
+3
View File
@@ -64,3 +64,6 @@ target/
# IntelliJ IDEA # IntelliJ IDEA
.idea/ .idea/
# Virtualenv
venv/
+1 -5
View File
@@ -6,8 +6,4 @@ max-attributes = 10
max-line-length=120 max-line-length=120
[MESSAGES CONTROL] [MESSAGES CONTROL]
disable = disable = missing-docstring, too-few-public-methods, locally-disabled, invalid-name
# Missing docstring
C0111,
# Too few public methods
R0903
+25
View File
@@ -23,6 +23,31 @@ client.send_notification(token_hex, payload)
[iOS Reference Library: Local and Push Notification Programming Guide][a1] [iOS Reference Library: Local and Push Notification Programming Guide][a1]
## Contributing
To develop PyAPNs2, check out the code and install dependencies. It's recommended to use a virtualenv to isolate dependencies:
```shell
# Clone the source code.
git clone https://github.com/Pr0ger/PyAPNs2.git
cd PyAPNs2
# Create a virtualenv and install dependencies.
virtualenv venv
. venv/bin/activate
pip install -e .
```
To run the tests:
```shell
pip install -r requirements-dev.txt
python -m unittest discover test
```
To run the linter:
```shell
pip install pylint
pylint --reports=n apns2 test
```
## License ## License
PyAPNs2 is distributed under the terms of the MIT license. PyAPNs2 is distributed under the terms of the MIT license.
+134 -17
View File
@@ -1,17 +1,27 @@
from enum import Enum import collections
from json import dumps
import json import json
import logging
from enum import Enum
from hyper import HTTP20Connection from hyper import HTTP20Connection
from hyper.tls import init_context from hyper.tls import init_context
from apns2.errors import exception_class_for_reason from .errors import ConnectionFailed, exception_class_for_reason
class NotificationPriority(Enum): class NotificationPriority(Enum):
Immediate = '10' Immediate = '10'
Delayed = '5' Delayed = '5'
RequestStream = collections.namedtuple('RequestStream', ['stream_id', 'token'])
Notification = collections.namedtuple('Notification', ['token', 'payload'])
DEFAULT_APNS_PRIORITY = NotificationPriority.Immediate
CONCURRENT_STREAMS_SAFETY_MAXIMUM = 1000
MAX_CONNECTION_RETRIES = 3
logger = logging.getLogger(__name__)
class APNsClient(object): class APNsClient(object):
def __init__(self, cert_file, use_sandbox=False, use_alternative_port=False, proto=None, json_encoder=None): def __init__(self, cert_file, use_sandbox=False, use_alternative_port=False, proto=None, json_encoder=None):
@@ -21,26 +31,133 @@ class APNsClient(object):
ssl_context.load_cert_chain(cert_file) ssl_context.load_cert_chain(cert_file)
self.__connection = HTTP20Connection(server, port, ssl_context=ssl_context, force_proto=proto or 'h2') self.__connection = HTTP20Connection(server, port, ssl_context=ssl_context, force_proto=proto or 'h2')
self.__json_encoder = json_encoder self.__json_encoder = json_encoder
self.__max_concurrent_streams = None
self.__previous_server_max_concurrent_streams = None
def send_notification(self, token_hex, notification, priority=NotificationPriority.Immediate, topic=None, def send_notification(self, token_hex, notification, topic, priority=NotificationPriority.Immediate,
expiration=None): expiration=None):
json_str = dumps(notification.dict(), cls=self.__json_encoder, ensure_ascii=False, separators=(',', ':')) stream_id = self.send_notification_async(token_hex, notification, topic, priority, expiration)
result = self.get_notification_result(stream_id)
if result != 'Success':
raise exception_class_for_reason(result)
def send_notification_async(self, token_hex, notification, topic, priority=NotificationPriority.Immediate,
expiration=None):
json_str = json.dumps(notification.dict(), cls=self.__json_encoder, ensure_ascii=False, separators=(',', ':'))
json_payload = json_str.encode('utf-8') json_payload = json_str.encode('utf-8')
headers = { headers = {'apns-topic': topic}
'apns-priority': priority.value if priority != DEFAULT_APNS_PRIORITY:
} headers['apns-priority'] = priority.value
if topic:
headers['apns-topic'] = topic
if expiration is not None: if expiration is not None:
headers['apns-expiration'] = "%d" % expiration headers['apns-expiration'] = '%d' % expiration
url = '/3/device/{}'.format(token_hex) 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)
resp = self.__connection.get_response(stream_id) return stream_id
with resp:
if resp.status != 200: def get_notification_result(self, stream_id):
raw_data = resp.read().decode('utf-8') with self.__connection.get_response(stream_id) as response:
if response.status == 200:
return 'Success'
else:
raw_data = response.read().decode('utf-8')
data = json.loads(raw_data) data = json.loads(raw_data)
raise exception_class_for_reason(data['reason']) return data['reason']
def send_notification_batch(self, notifications, topic, priority=NotificationPriority.Immediate, expiration=None):
'''
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,
using HTTP/2 streams (one request per stream).
APNs allows many streams simultaneously, but the number of streams can vary depending on
server load. This method reads the SETTINGS frame sent by the server to figure out the
maximum number of concurrent streams. Typically, APNs reports a maximum of 500.
The function returns a dictionary mapping each token to its result. The result is "Success"
if the token was sent successfully, or the string returned by APNs in the 'reason' field of
the response, if the token generated an error.
'''
notification_iterator = iter(notifications)
next_notification = next(notification_iterator, None)
# Make sure we're connected to APNs, so that we receive and process the server's SETTINGS
# frame before starting to send notifications.
self.connect()
results = {}
open_streams = collections.deque()
# 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.
while len(open_streams) > 0 or next_notification is not None:
# 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):
logger.info('Sending to token %s', next_notification.token)
stream_id = self.send_notification_async(next_notification.token, next_notification.payload, topic,
priority, expiration)
open_streams.append(RequestStream(stream_id, next_notification.token))
next_notification = next(notification_iterator, None)
if next_notification is None:
# No tokens remaining. Proceed to get results for pending requests.
logger.info('Finished sending all tokens, waiting for pending requests.')
else:
# We have at least one request waiting for response (otherwise we would have either
# sent new requests or exited the while loop.) Wait for the first outstanding stream
# to return a response.
pending_stream = open_streams.popleft()
result = self.get_notification_result(pending_stream.stream_id)
logger.info('Got response for %s: %s', pending_stream.token, result)
results[pending_stream.token] = result
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):
# 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.
# pylint: disable=protected-access
with self.__connection._conn as connection:
max_concurrent_streams = connection.remote_settings.max_concurrent_streams
if max_concurrent_streams == self.__previous_server_max_concurrent_streams:
# The server hasn't issued an updated SETTINGS frame.
return
self.__previous_server_max_concurrent_streams = max_concurrent_streams
# Handle and log unexpected values sent by APNs, just in case.
if max_concurrent_streams > CONCURRENT_STREAMS_SAFETY_MAXIMUM:
logger.warning('APNs max_concurrent_streams too high (%s), resorting to default maximum (%s)',
max_concurrent_streams, CONCURRENT_STREAMS_SAFETY_MAXIMUM)
self.__max_concurrent_streams = CONCURRENT_STREAMS_SAFETY_MAXIMUM
elif max_concurrent_streams < 1:
logger.warning('APNs reported max_concurrent_streams less than 1 (%s), using value of 1',
max_concurrent_streams)
self.__max_concurrent_streams = 1
else:
logger.info('APNs set max_concurrent_streams to %s', max_concurrent_streams)
self.__max_concurrent_streams = max_concurrent_streams
def connect(self):
'''
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:
try:
self.__connection.connect()
logger.info('Connected to APNs')
return
except Exception: # pylint: disable=broad-except
retries += 1
logger.exception('Failed connecting to APNs (attempt %s of %s)', retries, MAX_CONNECTION_RETRIES)
raise ConnectionFailed()
+4
View File
@@ -129,6 +129,10 @@ class MissingTopic(BadPayloadException):
pass pass
class ConnectionFailed(APNsException):
"""There was an error connecting to APNs."""
def exception_class_for_reason(reason): def exception_class_for_reason(reason):
return { return {
'PayloadEmpty': PayloadEmpty, 'PayloadEmpty': PayloadEmpty,
+1
View File
@@ -0,0 +1 @@
mock; python_version < '3.3'
+1 -1
View File
@@ -2,7 +2,7 @@
from distutils.core import setup from distutils.core import setup
dependencies = ['hyper'] dependencies = ['hyper>=0.7']
try: try:
# noinspection PyUnresolvedReferences # noinspection PyUnresolvedReferences
View File
+110
View File
@@ -0,0 +1,110 @@
# pylint: disable=protected-access
from unittest import TestCase
import contextlib
import logging
try:
# Python 3
from unittest.mock import MagicMock, Mock, patch
except ImportError:
# Python 2
from mock import MagicMock, Mock, patch
from apns2.client import APNsClient, CONCURRENT_STREAMS_SAFETY_MAXIMUM, Notification
from apns2.errors import ConnectionFailed
from apns2.payload import Payload
class ClientTestCase(TestCase):
@classmethod
def setUpClass(cls):
# Ignore all log messages so that test output is not cluttered.
logging.basicConfig(level=logging.CRITICAL)
cls.tokens = ['%064x' % i for i in range(10000)]
cls.payload = Payload(alert='Test alert')
cls.notifications = [Notification(token=token, payload=cls.payload) for token in cls.tokens]
cls.topic = 'com.example.App'
def setUp(self):
self.open_streams = 0
self.max_open_streams = 0
self.mock_results = None
self.next_stream_id = 0
with patch('apns2.client.HTTP20Connection') as mock_connection_constructor, patch('apns2.client.init_context'):
self.mock_connection = MagicMock()
self.mock_connection.get_response.side_effect = self.mock_get_response
self.mock_connection.request.side_effect = self.mock_request
self.mock_connection._conn.__enter__.return_value = self.mock_connection._conn
self.mock_connection._conn.remote_settings.max_concurrent_streams = 500
mock_connection_constructor.return_value = self.mock_connection
self.client = APNsClient(cert_file=None)
@contextlib.contextmanager
def mock_get_response(self, stream_id):
self.open_streams -= 1
if self.mock_results:
reason = self.mock_results[stream_id]
response = Mock(status=200 if reason == 'Success' else 400)
response.read.return_value = ('{"reason": "%s"}' % reason).encode('utf-8')
yield response
else:
yield Mock(status=200)
def mock_request(self, *dummy_args):
self.open_streams += 1
if self.open_streams > self.max_open_streams:
self.max_open_streams = self.open_streams
stream_id = self.next_stream_id
self.next_stream_id += 1
return stream_id
def test_send_notification_batch_returns_results_in_order(self):
results = self.client.send_notification_batch(self.notifications, self.topic)
expected_results = {token: 'Success' for token in self.tokens}
self.assertEqual(results, expected_results)
def test_send_notification_batch_respects_max_concurrent_streams_from_server(self):
self.client.send_notification_batch(self.notifications, self.topic)
self.assertEqual(self.max_open_streams, 500)
def test_send_notification_batch_overrides_server_max_concurrent_streams_if_too_large(self):
self.mock_connection._conn.remote_settings.max_concurrent_streams = 5000
self.client.send_notification_batch(self.notifications, self.topic)
self.assertEqual(self.max_open_streams, CONCURRENT_STREAMS_SAFETY_MAXIMUM)
def test_send_notification_batch_overrides_server_max_concurrent_streams_if_too_small(self):
self.mock_connection._conn.remote_settings.max_concurrent_streams = 0
self.client.send_notification_batch(self.notifications, self.topic)
self.assertEqual(self.max_open_streams, 1)
def test_send_notification_batch_reports_different_results(self):
self.mock_results = (
['BadDeviceToken'] * 1000 + ['Success'] * 1000 + ['DeviceTokenNotForTopic'] * 2000 +
['Success'] * 1000 + ['BadDeviceToken'] * 500 + ['PayloadTooLarge'] * 4500
)
results = self.client.send_notification_batch(self.notifications, self.topic)
expected_results = dict(zip(self.tokens, self.mock_results))
self.assertEqual(results, expected_results)
def test_send_empty_batch_does_nothing(self):
self.client.send_notification_batch([], self.topic)
self.assertEqual(self.mock_connection.request.call_count, 0)
def test_connect_establishes_connection(self):
self.client.connect()
self.mock_connection.connect.assert_called_once_with()
def test_connect_retries_failed_connection(self):
self.mock_connection.connect.side_effect = [RuntimeError, RuntimeError, None]
self.client.connect()
self.assertEqual(self.mock_connection.connect.call_count, 3)
def test_connect_stops_on_reaching_max_retries(self):
self.mock_connection.connect.side_effect = [RuntimeError] * 4
with self.assertRaises(ConnectionFailed):
self.client.connect()
self.assertEqual(self.mock_connection.connect.call_count, 3)