diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..4092704 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,29 @@ +name: CI + +on: + pull_request: + push: + branches: + - main + +jobs: + checks: + name: Compile, import, and tests + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + os: [ubuntu-latest, windows-latest, macos-latest] + python-version: ["3.8", "3.9", "3.10", "3.11"] + steps: + - uses: actions/checkout@v6 + - name: Set up Python + uses: actions/setup-python@v6 + with: + python-version: ${{ matrix.python-version }} + - name: Compile package + run: python -m compileall discordrpc + - name: Import package + run: python -c "import discordrpc" + - name: Run tests + run: python tests/run_tests.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 6456bc0..15c7ab3 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,32 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 --- +## [6.5b2] - Unreleased + +### Added +- Cross-platform event subscription with `Event` enum and `@rpc.on()` decorator (PR [#66](https://github.com/Senophyx/Discord-RPC/pull/66) by @SuperZombi) +- `InvalidEvent` and `InvalidEventType` exceptions exported from the package root +- `utils.required_url()` for validating required button URLs +- `examples/rpc-events.py` showing event subscription usage + +### Changed +- IPC reads now preserve the opcode and route responses by nonce through a single background reader +- Windows imports are now loaded only on Windows; the blocking reader no longer needs `msvcrt` or `win32pipe` +- URL validation now checks the scheme and host instead of only a prefix +- `set_activity()` now validates button URLs client-side + +### Fixed +- Invalid method annotations in `_send()` and `_request()` that prevented importing the package +- Unread `PONG` responses no longer corrupt subsequent RPC requests +- Incoming `PING` packets are answered with `PONG` +- Handshake no longer requires a nonce and reads the READY frame directly +- Event subscription state no longer mixes callback storage with subscription state +- Failed subscriptions no longer register callbacks locally +- Reader thread stops after the last unsubscribe and is cleaned up during disconnect +- `subscribe()` returns `True` for already-subscribed events so stacked handlers work +- Pipe is marked disconnected when the reader stops idle, allowing reconnect +- `pyproject.toml` license uses the PEP 621 table so wheel builds succeed + ## [Unreleased] ### Added diff --git a/DOCS.md b/DOCS.md index cf1f7c7..c85857c 100644 --- a/DOCS.md +++ b/DOCS.md @@ -295,6 +295,76 @@ from discordrpc import StatusDisplay --- +## Events + +You can subscribe to Rich Presence events and receive callbacks when they fire. + +```python +import discordrpc +from discordrpc import Event + +rpc = discordrpc.RPC(app_id=123456789) + +@rpc.on(Event.JOIN_REQUEST) +def on_join_request(data): + print("Ask to Join:", data) + +rpc.run() +``` + +Supported events (from `discordrpc.Event`): + +- `Event.JOIN` (`ACTIVITY_JOIN`) +- `Event.JOIN_REQUEST` (`ACTIVITY_JOIN_REQUEST`) +- `Event.SPECTATE` (`ACTIVITY_SPECTATE`) +- `Event.INVITE` (`ACTIVITY_INVITE`) + +Only these activity events are currently exposed. Other Discord RPC events are not supported yet. + +### Enabling JOIN and SPECTATE events + +- `party_id` is **required** for the "Ask to Join" button and the `ACTIVITY_JOIN_REQUEST` event to work. Without it, Discord cannot resolve the party, the event is never delivered, and the requester gets "Your message could not be delivered." +- `join_secret` and `spectate_secret` must have **different values**. Discord rejects the activity with `secrets must be unique` when they match. + +A minimal working setup: + +```python +rpc.set_activity( + name="VALORANT", + details="Valorant Ranked", + party_id=1234, + join_secret="anything", + spectate_secret="idk", +) +``` + +### Direct subscribe and unsubscribe + +```python +rpc.subscribe("ACTIVITY_JOIN") # string form accepted +rpc.subscribe(Event.JOIN) # enum form accepted +rpc.unsubscribe(Event.JOIN) +``` + +- `subscribe()` / `unsubscribe()` accept either an `Event` member or a valid event string. +- An unknown event name raises `InvalidEvent`. +- A non-string, non-enum value raises `InvalidEventType`. +- The `@rpc.on()` decorator raises `RPCException` if the subscription is rejected by Discord. + +### Callback behavior + +- Callbacks run on the internal IPC reader thread. Keep them short and non-blocking. +- A slow callback delays processing of other events and responses. +- An exception raised inside a callback is logged and does not stop the reader. +- Use locks or queues if your callback touches shared state. + +### Disconnect and reconnect + +- `disconnect()` clears all subscriptions and callbacks. +- After a reconnect, call `subscribe()` again if you want events. + +--- + ## Exceptions All exceptions extend `RPCException`. @@ -305,12 +375,14 @@ All exceptions extend `RPCException`. | `Error(message)` | Generic user error | | `DiscordNotOpened()` | Discord not found/running | | `ActivityError()` | Invalid activity payload | -| `InvalidURL()` | URL not starting with http/https | +| `InvalidURL(message)` | URL is not a valid http/https URL | | `InvalidID()` | Invalid Application ID | | `ButtonError(message)` | Button limit exceeded | | `ProgressbarError(message)` | Invalid progress values | | `InvalidActivityType(message)` | act_type not a valid Activity | | `ActivityTypeDisabled()` | Streaming/Custom blocked by Discord | +| `InvalidEvent(message)` | Event name is not subscribable | +| `InvalidEventType(message)` | Event input is not a string or Event | --- diff --git a/discordrpc/__init__.py b/discordrpc/__init__.py index 9fa0d96..f7ab8a6 100644 --- a/discordrpc/__init__.py +++ b/discordrpc/__init__.py @@ -5,6 +5,7 @@ RPCException, Error, DiscordNotOpened, ActivityError, InvalidURL, InvalidID, ButtonError, ProgressbarError, InvalidActivityType, ActivityTypeDisabled, + InvalidEvent, InvalidEventType, ) from .types import Activity, StatusDisplay, User, Application, Event from .utils import remove_none, timestamp, date_to_timestamp, use_local_time, progress_bar, get_app_info @@ -13,7 +14,7 @@ try: __version__ = _pkg_ver('discord-rpc') except PackageNotFoundError: - __version__ = "6.5b1" + __version__ = "6.5b2" __authors__ = "Senophyx" __license__ = "MIT License" __copyright__ = "Copyright 2021-2025 Senophyx" diff --git a/discordrpc/button.py b/discordrpc/button.py index 596380c..3c68cc6 100644 --- a/discordrpc/button.py +++ b/discordrpc/button.py @@ -1,6 +1,5 @@ -from .exceptions import InvalidURL -from .utils import valid_url +from .utils import required_url -def button(text:str, url:str): - return {"label": text, "url": valid_url(url)} +def button(text: str, url: str): + return {"label": text, "url": required_url(url)} diff --git a/discordrpc/exceptions.py b/discordrpc/exceptions.py index 232b5d1..cee7b4c 100644 --- a/discordrpc/exceptions.py +++ b/discordrpc/exceptions.py @@ -17,8 +17,10 @@ def __init__(self): super().__init__("An error has occurred in activity payload, do you have set your activity correctly?") class InvalidURL(RPCException): - def __init__(self): - super().__init__("URL must start with http:// or https://") + def __init__(self, message: str = None): + if message is None: + message = "URL must be a valid http:// or https:// URL" + super().__init__(message) class InvalidID(RPCException): def __init__(self): diff --git a/discordrpc/presence.py b/discordrpc/presence.py index 4b6cba3..1dd241f 100644 --- a/discordrpc/presence.py +++ b/discordrpc/presence.py @@ -4,6 +4,8 @@ import json import struct import uuid +import threading +import queue from typing import Optional from .exceptions import ( RPCException, InvalidID, DiscordNotOpened, @@ -11,15 +13,11 @@ InvalidEvent, InvalidEventType, ) from .types import Activity, StatusDisplay, User, Application, Asset, AssetManager, Event -from .utils import remove_none, get_app_info, get_assets, valid_url +from .utils import remove_none, get_app_info, get_assets, valid_url, required_url from functools import cached_property import logging import time -import msvcrt -import win32pipe -import threading - OP_HANDSHAKE = 0 OP_FRAME = 1 OP_CLOSE = 2 @@ -54,8 +52,9 @@ def __init__(self, app_id:int, debug:bool=False, output:bool=True, exit_if_disco log.disabled = True self.is_running = False - self._reader_thread = None self._event_callbacks = {} + self._subscriptions = set() + self._state_lock = threading.Lock() self._setup() def _setup(self): @@ -66,12 +65,7 @@ def _setup(self): if not self.ipc.connected: return self._user_data = self.ipc.handshake() - - def _start_event_listener(self): - if self._reader_thread and self._reader_thread.is_alive(): - return - self._reader_thread = threading.Thread(target=self._reader_loop, daemon=True) - self._reader_thread.start() + self.ipc._on_event = self._dispatch @property def connected(self): return self.ipc.connected @@ -119,6 +113,17 @@ def set_activity( if buttons and len(buttons) > 2: raise ButtonError("Max 2 buttons allowed") + # Validate button URLs so invalid URLs are caught client-side instead of + # being silently rejected by Discord (or missing from the presence). + if buttons: + buttons = [ + { + "label": item.get("label") if isinstance(item, dict) else None, + "url": required_url(item.get("url") if isinstance(item, dict) else item), + } + for item in buttons + ] + large_image = large_image.name if isinstance(large_image, Asset) else large_image small_image = small_image.name if isinstance(small_image, Asset) else small_image @@ -179,16 +184,14 @@ def set_activity( res = self.ipc._request(payload) if not res.get("ok"): self.is_running = False - log.error('Failed to set RPC') - log.error(res.get("error")) + log.error("Failed to set RPC: %s", res.get("error")) return False self.is_running = True - log.info('RPC set') + log.info("RPC set") return True - except Exception as e: - log.error('Failed to set RPC') - log.error(e) + except Exception: + log.exception("Failed to set RPC") self.disconnect() return False @@ -202,84 +205,100 @@ def disconnect(self): self.ipc.disconnect() self.is_running = False - - def subscribe(self, event: str): - if event not in [event.value for event in Event]: + with self._state_lock: + self._subscriptions.clear() + self._event_callbacks.clear() + + def _normalize_event(self, event): + if isinstance(event, Event): + return event.value + if isinstance(event, str): + values = {e.value for e in Event} + if event in values: + return event raise InvalidEvent(event) + raise InvalidEventType(type(event).__name__) + + def subscribe(self, event) -> Optional[bool]: + event = self._normalize_event(event) if not self.ipc.connected: return - if event in self._event_callbacks.keys(): - log.debug(f"Event {event} already registered") - return - + with self._state_lock: + if event in self._subscriptions: + log.debug(f"Event {event} already subscribed") + return True + payload = {"cmd": "SUBSCRIBE", "args": {}, "evt": event, "nonce": str(uuid.uuid4())} res = self.ipc._request(payload) if not res.get("ok"): - log.error(f'Failed to subscribe to {event}') - log.error(res.get("error")) + log.error("Failed to subscribe to %s: %s", event, res.get("error")) return False - self._event_callbacks.setdefault(event, []) + with self._state_lock: + self._subscriptions.add(event) + self._event_callbacks.setdefault(event, []) log.info(f"Subscribed to {event}") - self._start_event_listener() return True - def unsubscribe(self, event: str): - if event not in [event.value for event in Event]: - raise InvalidEvent(event) + def unsubscribe(self, event) -> Optional[bool]: + event = self._normalize_event(event) if not self.ipc.connected: return - if not event in self._event_callbacks.keys(): - log.error(f"Event {event} not registered") - return - + with self._state_lock: + if event not in self._subscriptions: + log.error(f"Event {event} not subscribed") + return + payload = {"cmd": "UNSUBSCRIBE", "args": {}, "evt": event, "nonce": str(uuid.uuid4())} res = self.ipc._request(payload) if not res.get("ok"): - log.error(f'Failed to unsubscribe from {event}') - log.error(res.get("error")) + log.error("Failed to unsubscribe from %s: %s", event, res.get("error")) return False - self._event_callbacks.pop(event, []) + with self._state_lock: + self._subscriptions.discard(event) + self._event_callbacks.pop(event, []) log.info(f"Unsubscribed from {event}") + self._stop_reader_if_idle() return True def on(self, event: Event): if type(event) != Event: - raise InvalidEventType(type(event)) - + raise InvalidEventType(type(event).__name__) + def decorator(callback): - self.subscribe(event.value) - self._event_callbacks.setdefault(event.value, []).append(callback) + result = self.subscribe(event) + if not result: + raise RPCException(f"Failed to subscribe to {event.value}") + with self._state_lock: + self._event_callbacks.setdefault(event.value, []).append(callback) return callback - + return decorator - def _reader_loop(self): - log.debug("Starting events listener") - while self.ipc.connected: + def _stop_reader_if_idle(self): + with self._state_lock: + idle = not self._subscriptions + if idle and self.ipc._reader_thread and self.ipc._reader_thread.is_alive(): + self.ipc._reader_stop.set() try: - frame = self.ipc.read_with_timeout() - except Exception as e: - log.debug(f"IPC reader stopping: {e}") - break - - if not frame: - continue - - if frame.get('cmd') == 'DISPATCH': - self._dispatch(frame.get('evt'), frame.get('data', {}) or {}) + self.ipc._close() + except Exception: + pass + self.ipc.connected = False def _dispatch(self, evt: str, data: dict): - for callback in self._event_callbacks.get(evt, []): + with self._state_lock: + callbacks = list(self._event_callbacks.get(evt, [])) + for callback in callbacks: try: callback(data) - except Exception as e: - log.error(f"Error in '{evt}' event callback: {e}") + except Exception: + log.exception("Error in '%s' event callback", evt) def run(self, update_every:int=1, ping_every:int=15): try: @@ -306,55 +325,196 @@ def __init__(self, app_id, exit_if_discord_close, exit_on_disconnect): self.exit_if_discord_close = exit_if_discord_close self.exit_on_disconnect = exit_on_disconnect self.connected = self._connect_pipe() + self._write_lock = threading.Lock() + self._pending_requests = {} + self._pending_lock = threading.Lock() + self._reader_thread = None + self._reader_stop = threading.Event() def _connect_pipe(self): """Override in subclass to establish the pipe connection. Returns True on success.""" raise NotImplementedError - def _send(self, payload, op=OP_FRAME: int): + def _send(self, payload, op: int = OP_FRAME): log.debug(payload) payload = json.dumps(payload).encode('UTF-8') payload = struct.pack(' dict: - self._send(payload, op) - res = self._recv() - if res.get("evt") == "ERROR": - return {"ok": False, "error": res.get("data", {}).get("message"), **res} - return {"ok": True, **res} + def _unregister_request(self, nonce): + with self._pending_lock: + self._pending_requests.pop(nonce, None) + + def _request(self, payload: dict, op: int = OP_FRAME, timeout: float = 10.0) -> dict: + nonce = payload.get("nonce") + if not nonce: + raise ValueError("RPC request payload must include a nonce") + + wait_queue = self._register_request(nonce) + self._start_reader() + try: + self._send(payload, op) + try: + res = wait_queue.get(timeout=timeout) + except queue.Empty: + return {"ok": False, "error": "RPC request timed out", "response": None} + finally: + self._unregister_request(nonce) + + opcode, payload_res = res + if opcode != OP_FRAME: + return {"ok": False, "error": f"Unexpected opcode {opcode}", "response": payload_res} + if not isinstance(payload_res, dict): + return {"ok": False, "error": "Expected JSON object response", "response": payload_res} + if payload_res.get("evt") == "ERROR": + data = payload_res.get("data") or {} + return { + "ok": False, + "code": data.get("code"), + "error": data.get("message") or payload_res.get("message"), + "response": payload_res, + } + if payload_res.get("nonce") != nonce: + return {"ok": False, "error": "RPC response nonce mismatch", "response": payload_res} + return {"ok": True, **payload_res} + + def _start_reader(self): + if self._reader_thread and self._reader_thread.is_alive(): + return + self._reader_stop.clear() + self._reader_thread = threading.Thread(target=self._reader_loop, daemon=True) + self._reader_thread.start() + + def _reader_loop(self): + while self.connected and not self._reader_stop.is_set(): + try: + frame = self._read_frame() + except Exception as e: + log.debug(f"IPC reader stopping: {e}") + break + if frame is None: + continue + opcode, payload = frame + if payload is None: + continue + if not isinstance(payload, dict): + continue + if opcode == OP_PING: + self._send(payload, OP_PONG) + continue + if opcode == OP_FRAME: + nonce = payload.get("nonce") + if nonce: + with self._pending_lock: + pending = self._pending_requests.get(nonce) + if pending: + pending.put((opcode, payload)) + continue + if payload.get("cmd") == "DISPATCH": + self._on_event(payload.get("evt"), payload.get("data") or {}) + continue + if opcode == OP_CLOSE: + log.debug("Received OP_CLOSE from Discord") + self.connected = False + break + + def _on_event(self, evt, data): + """Override in subclass or by RPC to dispatch events.""" def _write(self, data: bytes): """Override in subclass to write bytes to the pipe.""" raise NotImplementedError - def _recv(self): - """Override in subclass to receive data from the pipe.""" + def _read_some(self, size: int) -> bytes: + """Override in subclass to read up to size bytes from the pipe.""" raise NotImplementedError + def _read_exact(self, size: int) -> bytes: + chunks = [] + remaining = size + while remaining: + chunk = self._read_some(remaining) + if not chunk: + raise OSError("Connection closed while reading IPC frame") + chunks.append(chunk) + remaining -= len(chunk) + return b"".join(chunks) + + def _read_frame(self): + header = self._read_exact(8) + opcode, length = struct.unpack(" 16 * 1024 * 1024: + raise ValueError(f"Invalid IPC frame length: {length}") + if length == 0: + return opcode, None + payload_bytes = self._read_exact(length) + try: + payload = json.loads(payload_bytes.decode("UTF-8")) + except (UnicodeDecodeError, json.JSONDecodeError) as e: + raise ValueError(f"Invalid IPC frame payload: {e}") + return opcode, payload + + def _handle_handshake_error(self, payload): + code = payload.get("code") + message = payload.get("error") or "Handshake failed" + if code == 4000 or "invalid" in str(message).lower(): + raise InvalidID() + raise RPCException(f"Handshake failed: {message}") + def handshake(self): - data = self._request({'v': 1, 'client_id': self.app_id}, op=OP_HANDSHAKE) + self._send({'v': 1, 'client_id': self.app_id}, OP_HANDSHAKE) + + opcode, payload = self._read_frame() + if payload is None: + raise RPCException("Handshake did not receive a READY event") + + if opcode == OP_CLOSE: + self.connected = False + raise RPCException("Handshake closed by Discord before READY") + + if not isinstance(payload, dict): + raise RPCException("Handshake did not receive a READY event") + + if payload.get("evt") == "ERROR": + data = payload.get("data") or {} + error_payload = { + "ok": False, + "code": data.get("code"), + "error": data.get("message") or payload.get("message"), + "response": payload, + } + self._handle_handshake_error(error_payload) - if data.get('cmd') == 'DISPATCH' and data.get('evt') == 'READY': - user = data.get('data', {}).get('user') + if payload.get("cmd") == "DISPATCH" and payload.get("evt") == "READY": + user = payload.get("data", {}).get("user") if user: log.info(f"Connected to {user.get('username')} ({user.get('id')})") return user - if data.get('code') == 4000: - raise InvalidID() - - raise RPCException() + raise RPCException("Handshake did not receive a READY event") def disconnect(self): try: self._send({}, OP_CLOSE) self._close() except Exception as e: - log.debug("Socket closed before command was received") + log.debug("Socket closed before command was received: %s", e) + self._reader_stop.set() + self._close_pending() + reader = self._reader_thread + if reader and reader.is_alive() and reader is not threading.current_thread(): + reader.join(timeout=2) + self._reader_thread = None self.socket = None self.connected = False @@ -362,13 +522,17 @@ def disconnect(self): if self.exit_on_disconnect: sys.exit() + def _close_pending(self): + with self._pending_lock: + pending = self._pending_requests + self._pending_requests = {} + for queue in pending.values(): + queue.put((OP_CLOSE, None)) + def _close(self): """Override in subclass to close the socket.""" raise NotImplementedError - def read_with_timeout(self, timeout=1): - raise NotImplementedError - class WindowsPipe(_BasePipe): def _connect_pipe(self): @@ -405,46 +569,8 @@ def _write(self, data: bytes): def _close(self): self.socket.close() - def _recv(self): - enc_header = b'' - header_size = 8 - - while header_size: - enc_header += self.socket.read(header_size) - header_size -= len(enc_header) - - dec_header = struct.unpack("= header_size: - frame = self.socket.read(header_size) - if frame: - dec_header = struct.unpack(" bytes: + return self.socket.read(size) or b"" class UnixPipe(_BasePipe): @@ -480,35 +606,11 @@ def _connect_pipe(self): return True def _write(self, data: bytes): - self.socket.send(data) + self.socket.sendall(data) def _close(self): self.socket.shutdown(socket.SHUT_RDWR) self.socket.close() - def _recv(self): - enc_header = b'' - header_size = 8 - - while header_size: - chunk = self.socket.recv(header_size) - if not chunk: - break - enc_header += chunk - header_size -= len(chunk) - - dec_header = struct.unpack(" bytes: + return self.socket.recv(size) diff --git a/discordrpc/utils.py b/discordrpc/utils.py index e0abe95..d3f327b 100644 --- a/discordrpc/utils.py +++ b/discordrpc/utils.py @@ -1,6 +1,7 @@ import time import json import urllib.request +import urllib.parse from datetime import datetime import logging from .exceptions import ProgressbarError, InvalidURL @@ -30,11 +31,24 @@ def date_to_timestamp(date:str): datetime.strptime(date, "%d/%m/%Y-%H:%M:%S").timetuple() )) -def valid_url(url:str) -> str: - if url and not url.startswith(("http://", "https://")): - raise InvalidURL() +def _validate_url(url, required: bool): + if url is None: + if required: + raise InvalidURL("URL must be a valid http:// or https:// URL") + return url + if not isinstance(url, str) or not url: + raise InvalidURL("URL must be a valid http:// or https:// URL") + parsed = urllib.parse.urlparse(url) + if parsed.scheme not in ("http", "https") or not parsed.netloc: + raise InvalidURL("URL must be a valid http:// or https:// URL") return url +def valid_url(url): + return _validate_url(url, required=False) + +def required_url(url): + return _validate_url(url, required=True) + def use_local_time(): now = datetime.now() seconds_since_midnight = now.hour * 3600 + now.minute * 60 + now.second diff --git a/examples/rpc-events.py b/examples/rpc-events.py new file mode 100644 index 0000000..c948af3 --- /dev/null +++ b/examples/rpc-events.py @@ -0,0 +1,24 @@ +import discordrpc +from discordrpc import Event + +rpc = discordrpc.RPC(app_id=123456789) + +@rpc.on(Event.JOIN) +@rpc.on(Event.JOIN_REQUEST) +@rpc.on(Event.SPECTATE) +@rpc.on(Event.INVITE) +def on_event(data): + print(data) + +rpc.set_activity( + name="VALORANT", + details="Valorant Ranked", + party_id=1234, + join_secret="anything", + spectate_secret="idk" +) + +try: + rpc.run() +except KeyboardInterrupt: + rpc.disconnect() diff --git a/pyproject.toml b/pyproject.toml index f19c906..9d5454d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -5,12 +5,12 @@ build-backend = "setuptools.build_meta" [project] name = "discord-rpc" ##### VERSION ##### -version = "6.5b1" +version = "6.5b2" ################### description = "A Python wrapper for the Discord RPC API" readme = "README.md" requires-python = ">=3.7" -license = "MIT" +license = { file = "LICENSE" } authors = [{ name = "Senophyx", email = "contact@senophyx.id" }] keywords = ["Discord", "rpc", "discord rpc"] classifiers = [ diff --git a/tests/run_tests.py b/tests/run_tests.py new file mode 100644 index 0000000..4c38f00 --- /dev/null +++ b/tests/run_tests.py @@ -0,0 +1,12 @@ +import unittest +import sys +import os + +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))) + +loader = unittest.TestLoader() +suite = loader.discover(os.path.dirname(__file__), pattern="test_*.py") +result = unittest.TextTestRunner(verbosity=2).run(suite) + +if not result.wasSuccessful(): + sys.exit(1) diff --git a/tests/test_events.py b/tests/test_events.py new file mode 100644 index 0000000..0d91a10 --- /dev/null +++ b/tests/test_events.py @@ -0,0 +1,202 @@ +import json +import queue +import struct +import threading +import time +import unittest + +from discordrpc.exceptions import InvalidEvent, InvalidEventType +from discordrpc.presence import _BasePipe, OP_FRAME, RPC +from discordrpc.types import Event + + +class FakePipe(_BasePipe): + def __init__(self): + super().__init__(None, False, False) + self._in = queue.Queue() + self._out = queue.Queue() + self._buffer = b"" + self.connected = True + + def _connect_pipe(self): + return True + + def _write(self, data): + self._out.put(data) + # Auto-respond to any request carrying a nonce so subscribe()/unsubscribe() + # can complete without a real Discord server. + opcode, length = struct.unpack(" size: + self.responses.insert(0, data[size:]) + data = data[:size] + return data + + def _close(self): + pass + + +class HandshakeTests(unittest.TestCase): + def _pipe_with(self, frames): + pipe = FakePipe() + pipe.responses = [build_frame(*frame) for frame in frames] + return pipe + + def test_ready_event_succeeds(self): + pipe = self._pipe_with([ + (OP_FRAME, {"cmd": "DISPATCH", "evt": "READY", "data": {"user": {"id": "1", "username": "seno"}}}) + ]) + user = pipe.handshake() + self.assertEqual(user.get("username"), "seno") + # Handshake frame must be sent with opcode 0 and no nonce. + raw = pipe.sent[0] + opcode, length = struct.unpack(" size: + self.chunks.insert(0, chunk[size:]) + chunk = chunk[:size] + return chunk + + +def frame_bytes(opcode, payload): + data = json.dumps(payload).encode("UTF-8") + return struct.pack(" size: + self._in.put(data[size:]) + data = data[:size] + return data + + def _close(self): + self._close_called = True + + +def build_frame(opcode, payload): + import json + import struct + + data = json.dumps(payload).encode("UTF-8") + return struct.pack("