From 2f26d3af35553769b8a4b0fdcb25a6c6811d8d25 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8E=89=E6=BE=9C?= Date: Thu, 30 Jul 2026 12:20:52 +0800 Subject: [PATCH 1/2] fix: make stream lifecycle cancellation-safe --- dingtalk_stream/stream.py | 170 ++++++++++++++++++++++++++------------ tests/test_stream.py | 85 +++++++++++++++++++ 2 files changed, 202 insertions(+), 53 deletions(-) create mode 100644 tests/test_stream.py diff --git a/dingtalk_stream/stream.py b/dingtalk_stream/stream.py index 91dc243..74a5c50 100644 --- a/dingtalk_stream/stream.py +++ b/dingtalk_stream/stream.py @@ -27,6 +27,8 @@ class DingTalkStreamClient(object): OPEN_CONNECTION_API = DINGTALK_OPENAPI_ENDPOINT + '/v1.0/gateway/connections/open' TAG_DISCONNECT = 'disconnect' + HTTP_TIMEOUT_SECONDS = 10 + MAX_PENDING_TASKS = 100 def __init__(self, credential: Credential, logger: logging.Logger = None): self.credential: Credential = credential @@ -38,6 +40,9 @@ def __init__(self, credential: Credential, logger: logging.Logger = None): self._pre_started = False self._is_event_required = False self._access_token = {} + self._runner_task = None + self._stop_event = None + self._connection_tasks = set() def register_all_event_handler(self, handler: EventHandler): handler.dingtalk_client = self @@ -59,37 +64,91 @@ def pre_start(self): async def start(self): self.pre_start() + current_task = asyncio.current_task() + if self._runner_task is not None and not self._runner_task.done(): + raise RuntimeError('DingTalk stream client is already running') - while True: - try: - connection = self.open_connection() + self._runner_task = current_task + self._stop_event = asyncio.Event() + try: + while not self._stop_event.is_set(): + try: + loop = asyncio.get_running_loop() + connection = await loop.run_in_executor(None, self.open_connection) - if not connection: - self.logger.error('open connection failed') - await asyncio.sleep(10) - continue - self.logger.info('endpoint is %s', connection) + if self._stop_event.is_set(): + break + if not connection: + self.logger.error('open connection failed') + await self._wait_before_retry(10) + continue + self.logger.info('connecting to endpoint %s', connection['endpoint']) - uri = f'{connection["endpoint"]}?ticket={quote_plus(connection["ticket"])}' - async with websockets.connect(uri) as websocket: - self.websocket = websocket - asyncio.create_task(self.keepalive(websocket)) - async for raw_message in websocket: - json_message = json.loads(raw_message) - asyncio.create_task(self.background_task(json_message)) - except KeyboardInterrupt as e: - break - except (asyncio.exceptions.CancelledError, - websockets.exceptions.ConnectionClosedError) as e: - self.logger.error('[start] network exception, error=%s', e) - await asyncio.sleep(10) - continue - except Exception as e: - await asyncio.sleep(3) - self.logger.exception('unknown exception', e) - continue - finally: - pass + uri = f'{connection["endpoint"]}?ticket={quote_plus(connection["ticket"])}' + async with websockets.connect(uri) as websocket: + self.websocket = websocket + keepalive_task = asyncio.create_task(self.keepalive(websocket)) + try: + async for raw_message in websocket: + try: + json_message = json.loads(raw_message) + except (TypeError, json.JSONDecodeError): + self.logger.warning('invalid message, content=%r', raw_message) + continue + if len(self._connection_tasks) >= self.MAX_PENDING_TASKS: + await asyncio.wait( + self._connection_tasks, + return_when=asyncio.FIRST_COMPLETED, + ) + task = asyncio.create_task(self.background_task(json_message, websocket)) + self._connection_tasks.add(task) + task.add_done_callback(self._connection_tasks.discard) + finally: + keepalive_task.cancel() + await asyncio.gather(keepalive_task, return_exceptions=True) + await self._cancel_connection_tasks() + if self.websocket is websocket: + self.websocket = None + except asyncio.exceptions.CancelledError: + raise + except websockets.exceptions.ConnectionClosedError as e: + self.logger.error('[start] network exception, error=%s', e) + await self._wait_before_retry(10) + except Exception: + self.logger.exception('unknown exception') + await self._wait_before_retry(3) + finally: + await self._cancel_connection_tasks() + websocket = self.websocket + self.websocket = None + if websocket is not None: + await websocket.close() + self._runner_task = None + self._stop_event = None + + async def stop(self): + """Stop reconnecting and close the currently active websocket.""" + if self._stop_event is not None: + self._stop_event.set() + websocket = self.websocket + if websocket is not None: + await websocket.close() + + async def _wait_before_retry(self, delay): + if self._stop_event is None or self._stop_event.is_set(): + return + try: + await asyncio.wait_for(self._stop_event.wait(), timeout=delay) + except asyncio.TimeoutError: + pass + + async def _cancel_connection_tasks(self): + tasks = list(self._connection_tasks) + self._connection_tasks.clear() + for task in tasks: + task.cancel() + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) async def keepalive(self, ws, ping_interval=60): while True: @@ -99,15 +158,19 @@ async def keepalive(self, ws, ping_interval=60): except websockets.exceptions.ConnectionClosed: break - async def background_task(self, json_message): + async def background_task(self, json_message, websocket=None): + target_websocket = websocket if websocket is not None else self.websocket try: - route_result = await self.route_message(json_message) - if route_result == DingTalkStreamClient.TAG_DISCONNECT: - await self.websocket.close() - except Exception as e: - self.logger.error(f"error processing message: {e}") + route_result = await self.route_message(json_message, target_websocket) + if route_result == DingTalkStreamClient.TAG_DISCONNECT and target_websocket is not None: + await target_websocket.close() + except asyncio.exceptions.CancelledError: + raise + except Exception: + self.logger.exception('error processing message') - async def route_message(self, json_message): + async def route_message(self, json_message, websocket=None): + target_websocket = websocket if websocket is not None else self.websocket result = '' msg_type = json_message.get('type', '') ack = None @@ -132,18 +195,15 @@ async def route_message(self, json_message): json_message) else: self.logger.warning('unknown message, content=%s', json_message) - if ack: - await self.websocket.send(json.dumps(ack.to_dict())) + if ack and target_websocket is not None: + await target_websocket.send(json.dumps(ack.to_dict())) return result def start_forever(self): - while True: - try: - asyncio.run(self.start()) - except KeyboardInterrupt as e: - break - finally: - time.sleep(3) + try: + asyncio.run(self.start()) + except KeyboardInterrupt: + pass def open_connection(self): self.logger.info('open connection, url=%s' % DingTalkStreamClient.OPEN_CONNECTION_API) @@ -171,7 +231,8 @@ def open_connection(self): response_text = '' response = requests.post(DingTalkStreamClient.OPEN_CONNECTION_API, headers=request_headers, - data=request_body) + data=request_body, + timeout=self.HTTP_TIMEOUT_SECONDS) response_text = response.text response.raise_for_status() @@ -185,13 +246,12 @@ def get_host_ip(self): 查询本机ip地址 :return: ip """ - ip = "" + ip = '' try: - s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) - s.connect(('8.8.8.8', 80)) - ip = s.getsockname()[0] - finally: - s.close() + with socket.socket(socket.AF_INET, socket.SOCK_DGRAM) as sock: + sock.connect(('8.8.8.8', 80)) + return sock.getsockname()[0] + except OSError: return ip def reset_access_token(self): @@ -216,7 +276,8 @@ def get_access_token(self): response_text = '' response = requests.post(url, headers=request_headers, - data=json.dumps(values)) + data=json.dumps(values), + timeout=self.HTTP_TIMEOUT_SECONDS) response_text = response.text response.raise_for_status() @@ -243,7 +304,10 @@ def upload_to_dingtalk(self, image_content, filetype='image', filename='image.pn upload_url = f'https://oapi.dingtalk.com/media/upload?access_token={quote_plus(access_token)}' try: response_text = '' - response = requests.post(upload_url, data=values, files=files) + response = requests.post(upload_url, + data=values, + files=files, + timeout=self.HTTP_TIMEOUT_SECONDS) response_text = response.text if response.status_code == 401: self.reset_access_token() diff --git a/tests/test_stream.py b/tests/test_stream.py new file mode 100644 index 0000000..ab7dcf2 --- /dev/null +++ b/tests/test_stream.py @@ -0,0 +1,85 @@ +import asyncio +import json +import unittest +from unittest import mock + +from dingtalk_stream.credential import Credential +from dingtalk_stream.frames import AckMessage +from dingtalk_stream.stream import DingTalkStreamClient + + +class FakeWebSocket: + + def __init__(self): + self.closed = False + self.sent = [] + + async def close(self): + self.closed = True + + async def send(self, data): + self.sent.append(data) + + +class DingTalkStreamClientTest(unittest.IsolatedAsyncioTestCase): + + def setUp(self): + self.client = DingTalkStreamClient(Credential('client-id', 'client-secret')) + + async def test_start_propagates_cancellation(self): + self.client.open_connection = mock.Mock(return_value=None) + task = asyncio.create_task(self.client.start()) + await asyncio.sleep(0) + task.cancel() + + with self.assertRaises(asyncio.CancelledError): + await task + + self.assertIsNone(self.client._runner_task) + + async def test_stop_interrupts_retry_delay(self): + self.client.open_connection = mock.Mock(return_value=None) + task = asyncio.create_task(self.client.start()) + while self.client._stop_event is None: + await asyncio.sleep(0) + + await self.client.stop() + await asyncio.wait_for(task, timeout=1) + + self.assertIsNone(self.client._runner_task) + + async def test_ack_is_sent_to_source_websocket(self): + old_websocket = FakeWebSocket() + new_websocket = FakeWebSocket() + self.client.websocket = new_websocket + ack = AckMessage() + ack.code = 200 + ack.message = 'OK' + self.client.event_handler.raw_process = mock.AsyncMock(return_value=ack) + message = { + 'type': 'EVENT', + 'headers': {'topic': 'test-topic', 'messageId': 'message-id'}, + 'data': '{}', + } + + await self.client.route_message(message, old_websocket) + + self.assertEqual(1, len(old_websocket.sent)) + self.assertEqual(200, json.loads(old_websocket.sent[0])['code']) + self.assertEqual([], new_websocket.sent) + + def test_open_connection_has_timeout(self): + response = mock.Mock() + response.text = '{}' + response.json.return_value = {} + with mock.patch('dingtalk_stream.stream.requests.post', return_value=response) as post: + self.client.open_connection() + + self.assertEqual( + DingTalkStreamClient.HTTP_TIMEOUT_SECONDS, + post.call_args.kwargs['timeout'], + ) + + +if __name__ == '__main__': + unittest.main() From 9efbbf38ec4e4a8d1a1ff3322abba3e4efb67935 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8E=89=E6=BE=9C?= Date: Thu, 30 Jul 2026 16:34:35 +0800 Subject: [PATCH 2/2] fix: harden reconnect and message processing --- .gitattributes | 1 + .github/workflows/publish.yml | 45 ++++--- .github/workflows/test.yml | 52 ++++++++ README.md | 32 ++++- dingtalk_stream/frames.py | 19 --- dingtalk_stream/stream.py | 183 +++++++++++++++++++++++--- dingtalk_stream/version.py | 2 +- examples/calcbot/calcbot.py | 7 +- setup.py | 14 +- tests/test_stream.py | 238 ++++++++++++++++++++++++++++++++++ 10 files changed, 527 insertions(+), 66 deletions(-) create mode 100644 .gitattributes create mode 100644 .github/workflows/test.yml diff --git a/.gitattributes b/.gitattributes new file mode 100644 index 0000000..94919c5 --- /dev/null +++ b/.gitattributes @@ -0,0 +1 @@ +README.md -text whitespace=cr-at-eol diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml index cd9ea90..99a3c65 100644 --- a/.github/workflows/publish.yml +++ b/.github/workflows/publish.yml @@ -1,18 +1,27 @@ -name: Build and Publish Package -on: - release: - types: [published] -jobs: - build-and-publish: - runs-on: ubuntu-latest - steps: - - name: Checkout - uses: actions/checkout@v2 - - name: Build Package - run: | - python setup.py sdist bdist_wheel - - name: Publish Package - uses: pypa/gh-action-pypi-publish@master - with: - user: __token__ - password: ${{ secrets.PYPI_API_TOKEN }} +name: Build and Publish Package +on: + release: + types: [published] +jobs: + build-and-publish: + runs-on: ubuntu-latest + steps: + - name: Checkout + uses: actions/checkout@v4 + - name: Set up Python + uses: actions/setup-python@v5 + with: + python-version: '3.14' + - name: Install build tools + run: python -m pip install build + - name: Test + run: | + python -m pip install --editable . + python -m unittest discover -s tests -v + - name: Build Package + run: python -m build + - name: Publish Package + uses: pypa/gh-action-pypi-publish@release/v1 + with: + user: __token__ + password: ${{ secrets.PYPI_API_TOKEN }} diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml new file mode 100644 index 0000000..661595f --- /dev/null +++ b/.github/workflows/test.yml @@ -0,0 +1,52 @@ +name: Test + +on: + pull_request: + push: + branches: + - main + +permissions: + contents: read + +jobs: + test: + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + os: + - ubuntu-latest + python: + - '3.8' + - '3.9' + - '3.10' + - '3.11' + - '3.12' + - '3.13' + - '3.14' + include: + - os: windows-latest + python: '3.10' + - os: windows-latest + python: '3.14' + + steps: + - name: Checkout + uses: actions/checkout@v4 + + - name: Set up Python + uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python }} + cache: pip + cache-dependency-path: setup.py + + - name: Install + run: python -m pip install --editable . + + - name: Test + run: python -m unittest discover -s tests -v + + - name: Compile package and examples + run: python -m compileall -q dingtalk_stream tests examples diff --git a/README.md b/README.md index de686dd..ae18d57 100644 --- a/README.md +++ b/README.md @@ -68,7 +68,12 @@ class CalcBotHandler(dingtalk_stream.ChatbotHandler): incoming_message = dingtalk_stream.ChatbotMessage.from_dict(callback.data) expression = incoming_message.text.content.strip() try: - result = eval(expression) + operands = [float(part.strip()) for part in expression.split('+')] + if len(operands) < 2: + raise ValueError('only addition expressions are supported') + result = sum(operands) + if result.is_integer(): + result = int(result) except Exception as e: result = 'Error: %s' % e self.logger.info('%s = %s' % (expression, result)) @@ -97,14 +102,29 @@ if __name__ == '__main__': 有的时候,你需要在已有的 ioloop 中使用钉钉 Stream 模式,不使用 `start_forever` 方法。 -此时,可以使用 `client.start()` 代替 `client.start_forever()`。注意:需要在网络异常后重新启动 +此时,可以使用 `client.start()` 代替 `client.start_forever()`。`start()` 会在网络异常后自动重连, +调用 `stop()` 可以中断重连等待并关闭当前连接: ```Python +client_task = asyncio.create_task(client.start()) try: - await client.start() -except (asyncio.exceptions.CancelledError, - websockets.exceptions.ConnectionClosedError) as e: - ... # 处理网络断线异常 + await run_application() +finally: + await client.stop() + await client_task +``` + +可以通过 `websocket_connect_options` 透传 `websockets.connect()` 参数,例如调整握手和心跳超时: + +```Python +client = dingtalk_stream.DingTalkStreamClient( + credential, + websocket_connect_options={ + 'open_timeout': 10, + 'ping_interval': 20, + 'ping_timeout': 20, + }, +) ``` ## 开发教程 diff --git a/dingtalk_stream/frames.py b/dingtalk_stream/frames.py index 9a36a58..20c28ea 100644 --- a/dingtalk_stream/frames.py +++ b/dingtalk_stream/frames.py @@ -170,25 +170,6 @@ def __init__(self): self.data = {} self.extensions = {} - @classmethod - def from_dict(cls, d): - msg = SystemMessage() - data = '' - for name, value in d.items(): - if name == 'specVersion': - msg.spec_version = value - elif name == 'data': - data = value - elif name == 'type': - pass - elif name == 'headers': - msg.headers = Headers.from_dict(value) - else: - msg.extensions[name] = value - if data: - msg.data = json.loads(data) - return msg - def __str__(self): return 'SystemMessage(spec_version=%s, type=%s, headers=%s, data=%s, extensions=%s)' % ( self.spec_version, diff --git a/dingtalk_stream/stream.py b/dingtalk_stream/stream.py index 74a5c50..065eb32 100644 --- a/dingtalk_stream/stream.py +++ b/dingtalk_stream/stream.py @@ -2,9 +2,11 @@ import asyncio import asyncio.exceptions +from collections import OrderedDict import json import logging import platform +import random import time import requests import socket @@ -19,6 +21,7 @@ from .frames import SystemMessage from .frames import EventMessage from .frames import CallbackMessage +from .frames import AckMessage from .log import setup_default_logger from .utils import DINGTALK_OPENAPI_ENDPOINT from .version import VERSION_STRING @@ -29,8 +32,18 @@ class DingTalkStreamClient(object): TAG_DISCONNECT = 'disconnect' HTTP_TIMEOUT_SECONDS = 10 MAX_PENDING_TASKS = 100 + TASK_CANCELLATION_TIMEOUT_SECONDS = 5 + MAX_CACHED_MESSAGE_RESULTS = 10000 + MESSAGE_RESULT_TTL_SECONDS = 5 * 60 + RECONNECT_BASE_DELAY_SECONDS = 1 + RECONNECT_MAX_DELAY_SECONDS = 60 + RECONNECT_JITTER_SECONDS = 1 - def __init__(self, credential: Credential, logger: logging.Logger = None): + def __init__( + self, + credential: Credential, + logger: logging.Logger = None, + websocket_connect_options=None): self.credential: Credential = credential self.event_handler: EventHandler = EventHandler() self.callback_handler_map = {} @@ -43,6 +56,12 @@ def __init__(self, credential: Credential, logger: logging.Logger = None): self._runner_task = None self._stop_event = None self._connection_tasks = set() + self._orphaned_tasks = set() + self._inflight_messages = {} + self._message_results = OrderedDict() + self.websocket_connect_options = dict( + websocket_connect_options or {}, + ) def register_all_event_handler(self, handler: EventHandler): handler.dingtalk_client = self @@ -70,6 +89,7 @@ async def start(self): self._runner_task = current_task self._stop_event = asyncio.Event() + reconnect_attempt = 0 try: while not self._stop_event.is_set(): try: @@ -80,26 +100,32 @@ async def start(self): break if not connection: self.logger.error('open connection failed') - await self._wait_before_retry(10) + await self._wait_before_retry( + self._reconnect_delay(reconnect_attempt), + ) + reconnect_attempt += 1 continue self.logger.info('connecting to endpoint %s', connection['endpoint']) uri = f'{connection["endpoint"]}?ticket={quote_plus(connection["ticket"])}' - async with websockets.connect(uri) as websocket: + async with websockets.connect( + uri, + **self.websocket_connect_options) as websocket: self.websocket = websocket keepalive_task = asyncio.create_task(self.keepalive(websocket)) try: async for raw_message in websocket: + # Receiving any server frame proves that this + # connection is healthy; future reconnects + # start from the base delay again. + reconnect_attempt = 0 try: json_message = json.loads(raw_message) except (TypeError, json.JSONDecodeError): self.logger.warning('invalid message, content=%r', raw_message) continue - if len(self._connection_tasks) >= self.MAX_PENDING_TASKS: - await asyncio.wait( - self._connection_tasks, - return_when=asyncio.FIRST_COMPLETED, - ) + if not await self._wait_for_task_capacity(): + break task = asyncio.create_task(self.background_task(json_message, websocket)) self._connection_tasks.add(task) task.add_done_callback(self._connection_tasks.discard) @@ -109,14 +135,25 @@ async def start(self): await self._cancel_connection_tasks() if self.websocket is websocket: self.websocket = None + if not self._stop_event.is_set(): + await self._wait_before_retry( + self._reconnect_delay(reconnect_attempt), + ) + reconnect_attempt += 1 except asyncio.exceptions.CancelledError: raise except websockets.exceptions.ConnectionClosedError as e: self.logger.error('[start] network exception, error=%s', e) - await self._wait_before_retry(10) + await self._wait_before_retry( + self._reconnect_delay(reconnect_attempt), + ) + reconnect_attempt += 1 except Exception: self.logger.exception('unknown exception') - await self._wait_before_retry(3) + await self._wait_before_retry( + self._reconnect_delay(reconnect_attempt), + ) + reconnect_attempt += 1 finally: await self._cancel_connection_tasks() websocket = self.websocket @@ -142,13 +179,54 @@ async def _wait_before_retry(self, delay): except asyncio.TimeoutError: pass + def _reconnect_delay(self, attempt): + exponential_delay = ( + self.RECONNECT_BASE_DELAY_SECONDS * (2 ** min(attempt, 16)) + ) + jitter = random.uniform(0, self.RECONNECT_JITTER_SECONDS) + return min( + exponential_delay + jitter, + self.RECONNECT_MAX_DELAY_SECONDS, + ) + async def _cancel_connection_tasks(self): tasks = list(self._connection_tasks) self._connection_tasks.clear() for task in tasks: task.cancel() if tasks: - await asyncio.gather(*tasks, return_exceptions=True) + _, pending = await asyncio.wait( + tasks, + timeout=self.TASK_CANCELLATION_TIMEOUT_SECONDS, + ) + if pending: + self.logger.warning( + '%d background task(s) ignored cancellation; ' + 'they remain counted against the global task limit', + len(pending), + ) + self._orphaned_tasks.update(pending) + for task in pending: + task.add_done_callback(self._orphaned_tasks.discard) + + async def _wait_for_task_capacity(self): + pending_tasks = self._connection_tasks | self._orphaned_tasks + while len(pending_tasks) >= self.MAX_PENDING_TASKS: + stop_event = self._stop_event + if stop_event is None or stop_event.is_set(): + return False + + stop_waiter = asyncio.create_task(stop_event.wait()) + try: + await asyncio.wait( + [*pending_tasks, stop_waiter], + return_when=asyncio.FIRST_COMPLETED, + ) + finally: + stop_waiter.cancel() + await asyncio.gather(stop_waiter, return_exceptions=True) + pending_tasks = self._connection_tasks | self._orphaned_tasks + return self._stop_event is not None and not self._stop_event.is_set() async def keepalive(self, ws, ping_interval=60): while True: @@ -161,16 +239,70 @@ async def keepalive(self, ws, ping_interval=60): async def background_task(self, json_message, websocket=None): target_websocket = websocket if websocket is not None else self.websocket try: - route_result = await self.route_message(json_message, target_websocket) - if route_result == DingTalkStreamClient.TAG_DISCONNECT and target_websocket is not None: - await target_websocket.close() + await self._background_task(json_message, target_websocket) except asyncio.exceptions.CancelledError: raise except Exception: self.logger.exception('error processing message') + async def _background_task(self, json_message, target_websocket): + message_type = json_message.get('type', '') + if message_type in (EventMessage.TYPE, CallbackMessage.TYPE): + message_id = json_message.get('headers', {}).get('messageId') + else: + # SYSTEM commands belong to one connection lifecycle. Replaying a + # cached disconnect result on a replacement connection could close + # the healthy socket, so control frames must always be processed. + message_id = None + if not message_id: + route_result, ack = await self._dispatch_message(json_message) + await self._send_ack(ack, target_websocket) + if route_result == DingTalkStreamClient.TAG_DISCONNECT and target_websocket is not None: + await target_websocket.close() + return + + cached = self._get_cached_message_result(message_id) + if cached is not None: + route_result, ack = cached + await self._send_ack(ack, target_websocket) + if route_result == DingTalkStreamClient.TAG_DISCONNECT and target_websocket is not None: + await target_websocket.close() + return + + inflight = self._inflight_messages.get(message_id) + if inflight is not None: + route_result, ack = await asyncio.shield(inflight) + await self._send_ack(ack, target_websocket) + if route_result == DingTalkStreamClient.TAG_DISCONNECT and target_websocket is not None: + await target_websocket.close() + return + + loop = asyncio.get_running_loop() + result_future = loop.create_future() + self._inflight_messages[message_id] = result_future + try: + route_result, ack = await self._dispatch_message(json_message) + if ack is not None and ack.code == AckMessage.STATUS_OK: + self._cache_message_result(message_id, route_result, ack) + result_future.set_result((route_result, ack)) + await self._send_ack(ack, target_websocket) + if route_result == DingTalkStreamClient.TAG_DISCONNECT and target_websocket is not None: + await target_websocket.close() + except BaseException: + if not result_future.done(): + result_future.cancel() + raise + finally: + if self._inflight_messages.get(message_id) is result_future: + self._inflight_messages.pop(message_id, None) + async def route_message(self, json_message, websocket=None): target_websocket = websocket if websocket is not None else self.websocket + result, ack = await self._dispatch_message(json_message) + await self._send_ack(ack, target_websocket) + return result + + async def _dispatch_message(self, json_message): result = '' msg_type = json_message.get('type', '') ack = None @@ -180,8 +312,6 @@ async def route_message(self, json_message, websocket=None): if msg.headers.topic == SystemMessage.TOPIC_DISCONNECT: result = DingTalkStreamClient.TAG_DISCONNECT self.logger.info("received disconnect topic=%s, message=%s", msg.headers.topic, json_message) - else: - self.logger.warning("unknown message topic, topic=%s, message=%s", msg.headers.topic, json_message) elif msg_type == EventMessage.TYPE: msg = EventMessage.from_dict(json_message) ack = await self.event_handler.raw_process(msg) @@ -195,9 +325,28 @@ async def route_message(self, json_message, websocket=None): json_message) else: self.logger.warning('unknown message, content=%s', json_message) + return result, ack + + async def _send_ack(self, ack, target_websocket): if ack and target_websocket is not None: await target_websocket.send(json.dumps(ack.to_dict())) - return result + + def _get_cached_message_result(self, message_id): + cached = self._message_results.get(message_id) + if cached is None: + return None + cached_at, route_result, ack = cached + if time.monotonic() - cached_at > self.MESSAGE_RESULT_TTL_SECONDS: + self._message_results.pop(message_id, None) + return None + self._message_results.move_to_end(message_id) + return route_result, ack + + def _cache_message_result(self, message_id, route_result, ack): + self._message_results[message_id] = (time.monotonic(), route_result, ack) + self._message_results.move_to_end(message_id) + while len(self._message_results) > self.MAX_CACHED_MESSAGE_RESULTS: + self._message_results.popitem(last=False) def start_forever(self): try: diff --git a/dingtalk_stream/version.py b/dingtalk_stream/version.py index d14dfec..fd929ef 100644 --- a/dingtalk_stream/version.py +++ b/dingtalk_stream/version.py @@ -1 +1 @@ -VERSION_STRING = '0.24.3' +VERSION_STRING = '0.24.4' diff --git a/examples/calcbot/calcbot.py b/examples/calcbot/calcbot.py index 755cc9d..be4d69f 100644 --- a/examples/calcbot/calcbot.py +++ b/examples/calcbot/calcbot.py @@ -39,7 +39,12 @@ async def process(self, callback: dingtalk_stream.CallbackMessage): incoming_message = dingtalk_stream.ChatbotMessage.from_dict(callback.data) expression = incoming_message.text.content.strip() try: - result = eval(expression) + operands = [float(part.strip()) for part in expression.split('+')] + if len(operands) < 2: + raise ValueError('only addition expressions are supported') + result = sum(operands) + if result.is_integer(): + result = int(result) except Exception as e: result = 'Error: %s' % e self.logger.info('%s = %s' % (expression, result)) diff --git a/setup.py b/setup.py index 9defe59..855ee85 100644 --- a/setup.py +++ b/setup.py @@ -4,21 +4,24 @@ BASE_PATH = os.path.dirname(os.path.abspath(__file__)) VERSION_STRING = '' -with open(os.path.join(BASE_PATH, 'dingtalk_stream', 'version.py')) as fp: +with open(os.path.join(BASE_PATH, 'dingtalk_stream', 'version.py'), encoding='utf-8') as fp: content = fp.read() VERSION_STRING = re.findall(r"VERSION_STRING\s*=\s*\'(.*?)\'", content)[0] +with open(os.path.join(BASE_PATH, 'README.md'), encoding='utf-8') as fp: + LONG_DESCRIPTION = fp.read() setup( name='dingtalk-stream', version=VERSION_STRING, description='A Python library for sending messages to DingTalk chatbot', - long_description=open('README.md').read(), + long_description=LONG_DESCRIPTION, long_description_content_type='text/markdown', url='https://github.com/open-dingtalk/dingtalk-stream-sdk-python', author='Ke Jie', author_email='jinxi.kj@alibaba-inc.com', license='MIT', packages=['dingtalk_stream'], + python_requires='>=3.8', install_requires=[ 'websockets>=11.0.2', 'requests>=2.27.1', @@ -29,10 +32,13 @@ 'Intended Audience :: Developers', 'License :: OSI Approved :: MIT License', 'Programming Language :: Python :: 3', - 'Programming Language :: Python :: 3.6', - 'Programming Language :: Python :: 3.7', 'Programming Language :: Python :: 3.8', 'Programming Language :: Python :: 3.9', + 'Programming Language :: Python :: 3.10', + 'Programming Language :: Python :: 3.11', + 'Programming Language :: Python :: 3.12', + 'Programming Language :: Python :: 3.13', + 'Programming Language :: Python :: 3.14', 'Programming Language :: Python :: Implementation :: CPython', 'Programming Language :: Python :: Implementation :: PyPy' ], diff --git a/tests/test_stream.py b/tests/test_stream.py index ab7dcf2..217152c 100644 --- a/tests/test_stream.py +++ b/tests/test_stream.py @@ -21,6 +21,27 @@ async def send(self, data): self.sent.append(data) +class EmptyWebSocket(FakeWebSocket): + + def __aiter__(self): + return self + + async def __anext__(self): + raise StopAsyncIteration + + +class FakeWebSocketConnect: + + def __init__(self, websocket): + self.websocket = websocket + + async def __aenter__(self): + return self.websocket + + async def __aexit__(self, exc_type, exc_value, traceback): + return False + + class DingTalkStreamClientTest(unittest.IsolatedAsyncioTestCase): def setUp(self): @@ -48,6 +69,115 @@ async def test_stop_interrupts_retry_delay(self): self.assertIsNone(self.client._runner_task) + def test_reconnect_delay_uses_bounded_exponential_backoff(self): + with mock.patch( + 'dingtalk_stream.stream.random.uniform', + return_value=0.5): + self.assertEqual(1.5, self.client._reconnect_delay(0)) + self.assertEqual(2.5, self.client._reconnect_delay(1)) + self.assertEqual(4.5, self.client._reconnect_delay(2)) + self.assertEqual( + self.client.RECONNECT_MAX_DELAY_SECONDS, + self.client._reconnect_delay(100), + ) + + async def test_normal_websocket_close_uses_reconnect_backoff(self): + options = { + 'open_timeout': 3, + 'ping_interval': None, + } + self.client = DingTalkStreamClient( + Credential('client-id', 'client-secret'), + websocket_connect_options=options, + ) + options['open_timeout'] = 99 + self.client.open_connection = mock.Mock(return_value={ + 'endpoint': 'ws://localhost', + 'ticket': 'test-ticket', + }) + observed_delays = [] + + async def stop_after_delay(delay): + observed_delays.append(delay) + self.client._stop_event.set() + + self.client._reconnect_delay = mock.Mock(return_value=1.25) + self.client._wait_before_retry = stop_after_delay + with mock.patch( + 'dingtalk_stream.stream.websockets.connect', + return_value=FakeWebSocketConnect(EmptyWebSocket())) as connect: + await self.client.start() + + self.assertEqual([1.25], observed_delays) + self.client._reconnect_delay.assert_called_once_with(0) + self.assertEqual(1, self.client.open_connection.call_count) + connect.assert_called_once_with( + 'ws://localhost?ticket=test-ticket', + open_timeout=3, + ping_interval=None, + ) + + async def test_stop_interrupts_wait_for_background_task_capacity(self): + self.client._stop_event = asyncio.Event() + blocker = asyncio.Event() + tasks = { + asyncio.create_task(blocker.wait()) + for _ in range(self.client.MAX_PENDING_TASKS) + } + self.client._connection_tasks.update(tasks) + capacity_waiter = asyncio.create_task(self.client._wait_for_task_capacity()) + await asyncio.sleep(0) + + await self.client.stop() + + self.assertFalse(await asyncio.wait_for(capacity_waiter, timeout=1)) + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + self.client._connection_tasks.clear() + + async def test_task_ignoring_cancellation_remains_globally_bounded(self): + self.client.TASK_CANCELLATION_TIMEOUT_SECONDS = 0.01 + cancellation_received = asyncio.Event() + release_task = asyncio.Event() + + async def ignore_cancellation(): + try: + await asyncio.Future() + except asyncio.CancelledError: + cancellation_received.set() + await release_task.wait() + + task = asyncio.create_task(ignore_cancellation()) + self.client._connection_tasks.add(task) + await asyncio.sleep(0) + + await asyncio.wait_for(self.client._cancel_connection_tasks(), timeout=1) + + self.assertTrue(cancellation_received.is_set()) + self.assertNotIn(task, self.client._connection_tasks) + self.assertIn(task, self.client._orphaned_tasks) + + self.client._stop_event = asyncio.Event() + blockers = { + asyncio.create_task(asyncio.sleep(3600)) + for _ in range(self.client.MAX_PENDING_TASKS - 1) + } + self.client._connection_tasks.update(blockers) + capacity_waiter = asyncio.create_task(self.client._wait_for_task_capacity()) + await asyncio.sleep(0) + self.assertFalse(capacity_waiter.done()) + + release_task.set() + await task + await asyncio.wait_for(capacity_waiter, timeout=1) + self.assertNotIn(task, self.client._orphaned_tasks) + + for blocker in blockers: + blocker.cancel() + await asyncio.gather(*blockers, return_exceptions=True) + self.client._connection_tasks.clear() + async def test_ack_is_sent_to_source_websocket(self): old_websocket = FakeWebSocket() new_websocket = FakeWebSocket() @@ -68,6 +198,114 @@ async def test_ack_is_sent_to_source_websocket(self): self.assertEqual(200, json.loads(old_websocket.sent[0])['code']) self.assertEqual([], new_websocket.sent) + async def test_duplicate_message_shares_handler_and_replays_ack(self): + first_websocket = FakeWebSocket() + retry_websocket = FakeWebSocket() + cached_retry_websocket = FakeWebSocket() + handler_started = asyncio.Event() + release_handler = asyncio.Event() + handler_calls = 0 + + async def process_once(_): + nonlocal handler_calls + handler_calls += 1 + handler_started.set() + await release_handler.wait() + ack = AckMessage() + ack.code = 200 + ack.message = 'OK' + return ack + + self.client.event_handler.raw_process = process_once + message = { + 'type': 'EVENT', + 'headers': {'topic': 'test-topic', 'messageId': 'duplicate-message-id'}, + 'data': '{}', + } + + first = asyncio.create_task(self.client.background_task(message, first_websocket)) + await handler_started.wait() + retry = asyncio.create_task(self.client.background_task(message, retry_websocket)) + await asyncio.sleep(0) + release_handler.set() + await asyncio.gather(first, retry) + await self.client.background_task(message, cached_retry_websocket) + + self.assertEqual(1, handler_calls) + self.assertEqual(1, len(first_websocket.sent)) + self.assertEqual(1, len(retry_websocket.sent)) + self.assertEqual(1, len(cached_retry_websocket.sent)) + self.assertEqual({}, self.client._inflight_messages) + + async def test_failed_message_result_is_not_cached(self): + handler_calls = 0 + + async def fail_for_retry(_): + nonlocal handler_calls + handler_calls += 1 + ack = AckMessage() + ack.code = AckMessage.STATUS_SYSTEM_EXCEPTION + ack.message = 'retry' + return ack + + self.client.event_handler.raw_process = fail_for_retry + message = { + 'type': 'EVENT', + 'headers': {'topic': 'test-topic', 'messageId': 'retry-message-id'}, + 'data': '{}', + } + + await self.client.background_task(message, FakeWebSocket()) + await self.client.background_task(message, FakeWebSocket()) + + self.assertEqual(2, handler_calls) + self.assertNotIn('retry-message-id', self.client._message_results) + + async def test_system_message_is_not_deduplicated_across_connections(self): + ack = AckMessage() + ack.code = AckMessage.STATUS_OK + ack.message = 'OK' + self.client.logger = mock.Mock() + self.client.system_handler.raw_process = mock.AsyncMock(return_value=ack) + message = { + 'type': 'SYSTEM', + 'headers': {'topic': 'ping', 'messageId': 'system-message-id'}, + 'data': '{}', + } + + await self.client.background_task(message, FakeWebSocket()) + await self.client.background_task(message, FakeWebSocket()) + + self.assertEqual(2, self.client.system_handler.raw_process.await_count) + self.assertNotIn('system-message-id', self.client._message_results) + self.client.logger.warning.assert_not_called() + + def test_message_result_cache_is_bounded_and_expires(self): + self.client.MAX_CACHED_MESSAGE_RESULTS = 3 + ack = AckMessage() + ack.code = AckMessage.STATUS_OK + ack.message = 'OK' + + with mock.patch('dingtalk_stream.stream.time.monotonic', return_value=100): + for index in range(5): + self.client._cache_message_result( + 'message-%d' % index, + '', + ack, + ) + + self.assertEqual( + ['message-2', 'message-3', 'message-4'], + list(self.client._message_results), + ) + with mock.patch( + 'dingtalk_stream.stream.time.monotonic', + return_value=100 + self.client.MESSAGE_RESULT_TTL_SECONDS + 1): + self.assertIsNone( + self.client._get_cached_message_result('message-4'), + ) + self.assertNotIn('message-4', self.client._message_results) + def test_open_connection_has_timeout(self): response = mock.Mock() response.text = '{}'