From 532ba52500d482d80f69a80c89df59a19a13702d Mon Sep 17 00:00:00 2001 From: Choco <94597336+ChocoMeow@users.noreply.github.com> Date: Fri, 16 Feb 2024 10:38:40 +0800 Subject: [PATCH] Optimized code --- voicelink/__init__.py | 1 + voicelink/pool.py | 135 ++++++++++++++++++----------------- voicelink/spotify/objects.py | 2 +- 3 files changed, 71 insertions(+), 67 deletions(-) diff --git a/voicelink/__init__.py b/voicelink/__init__.py index d48084f..75673a8 100644 --- a/voicelink/__init__.py +++ b/voicelink/__init__.py @@ -36,3 +36,4 @@ from .pool import * from .queue import * from .placeholders import Placeholders, build_embed from .formatter import encode, decode +from .spotify import * diff --git a/voicelink/pool.py b/voicelink/pool.py index 6b33e2a..7a7fc6c 100644 --- a/voicelink/pool.py +++ b/voicelink/pool.py @@ -29,6 +29,7 @@ import re import aiohttp from discord import Client, Member +from discord.ext.commands import Bot from typing import Dict, Optional, TYPE_CHECKING, Union from urllib.parse import quote @@ -69,10 +70,6 @@ URL_REGEX = re.compile( NODE_VERSION = "v4" CALL_METHOD = ["PATCH", "DELETE"] -def exception_catch_callback(task): - if task.exception(): - return - class Node: """The base class for a node. This node object represents a Lavalink node. @@ -83,7 +80,7 @@ class Node: self, *, pool, - bot: Client, + bot: Bot, host: str, port: int, password: str, @@ -96,29 +93,27 @@ class Node: resume_key: Optional[str] = None ): - self._bot = bot - self._host = host - self._port = port - self._pool = pool - self._password = password - self._identifier = identifier - self._heartbeat = heartbeat - self._secure = secure + self._bot: Bot = bot + self._host: str = host + self._port: int = port + self._pool: NodePool = pool + self._password: str = password + self._identifier: str = identifier + self._heartbeat: int = heartbeat + self._secure: bool = secure - self._websocket_uri = f"{'wss' if self._secure else 'ws'}://{self._host}:{self._port}/" + NODE_VERSION + "/websocket" - self._rest_uri = f"{'https' if self._secure else 'http'}://{self._host}:{self._port}" + self._websocket_uri: str = f"{'wss' if self._secure else 'ws'}://{self._host}:{self._port}/" + NODE_VERSION + "/websocket" + self._rest_uri: str = f"{'https' if self._secure else 'http'}://{self._host}:{self._port}" - self._session = session or aiohttp.ClientSession() + self._session: aiohttp.ClientSession = session or aiohttp.ClientSession() self._websocket: aiohttp.ClientWebSocketResponse = None self._task: asyncio.Task = None - self.resume_key = resume_key or str(os.urandom(8).hex()) + self.resume_key: str = resume_key or str(os.urandom(8).hex()) + self._session_id: str = None + self._available: bool = None - self._session_id = None - self._metadata = None - self._available = None - - self._headers = { + self._headers: Dict[str, str] = { "Authorization": self._password, "User-Id": str(bot.user.id), "Client-Name": f"Voicelink/{__version__}", @@ -127,14 +122,10 @@ class Node: self._players: Dict[int, Player] = {} - self._spotify_client_id = spotify_client_id - self._spotify_client_secret = spotify_client_secret - - if self._spotify_client_id and self._spotify_client_secret: - self._spotify_client = spotify.Client( - self._spotify_client_id, self._spotify_client_secret - ) - + self._spotify_client_id: Optional[str] = spotify_client_id + self._spotify_client_secret: Optional[str] = spotify_client_secret + self._spotify_client: Optional[spotify.Client] = None + self._bot.add_listener(self._update_handler, "on_socket_response") def __repr__(self): @@ -142,7 +133,19 @@ class Node: f"" ) + + @property + def spotify_client(self) -> Optional[spotify.Client]: + if not self._spotify_client: + if not self._spotify_client_id and not self._spotify_client_secret: + return None + + self._spotify_client = spotify.Client( + self._spotify_client_id, self._spotify_client_secret + ) + return self._spotify_client + @property def is_connected(self) -> bool: """"Property which returns whether this node is connected or not""" @@ -160,7 +163,7 @@ class Node: return self._players @property - def bot(self) -> Client: + def bot(self) -> Bot: """Property which returns the discord.py client linked to this node""" return self._bot @@ -170,16 +173,16 @@ class Node: return len(self.players) @property - def pool(self): + def pool(self) -> NodePool: """Property which returns the pool this node is apart of""" return self._pool @property - def latency(self): + def latency(self) -> float: """Property which returns the latency of the node""" return Ping(self._host, port=self._port).get_ping() - async def _update_handler(self, data: dict): + async def _update_handler(self, data: dict) -> None: #await self._bot.wait_until_ready() if not data: @@ -204,7 +207,7 @@ class Node: except KeyError: return - async def _listen(self): + async def _listen(self) -> None: backoff = ExponentialBackoff(base=7) while True: @@ -226,7 +229,7 @@ class Node: else: self._bot.loop.create_task(self._handle_payload(msg.json())) - async def _handle_payload(self, data: dict): + async def _handle_payload(self, data: dict) -> None: op = data.get("op", None) if not op: return @@ -247,25 +250,26 @@ class Node: elif op == "playerUpdate": await player._update_state(data) - async def send(self, method: int, - guild_id: Union[str, int] = None, - query: str = None, - data: Union[dict, str] = {}): - + async def send( + self, method: int, + guild_id: Union[str, int] = None, + query: str = None, + data: Union[dict, str] = {} + ) -> dict: if not self._available: - raise NodeNotAvailable( - f"The node '{self._identifier}' is unavailable." - ) + raise NodeNotAvailable(f"The node '{self._identifier}' is unavailable.") - uri: str = f"{self._rest_uri}/{NODE_VERSION}" \ - f"/sessions/{self._session_id}/players" \ - f"/{guild_id}" if guild_id else "" \ - f"?{query}" if query else "" + uri: str = f"{self._rest_uri}/{NODE_VERSION}" \ + f"/sessions/{self._session_id}/players" \ + f"/{guild_id}" if guild_id else "" \ + f"?{query}" if query else "" - async with self._session.request(method=CALL_METHOD[method], - url=uri, - headers={"Authorization": self._password}, - json=data) as resp: + async with self._session.request( + method=CALL_METHOD[method], + url=uri, + headers={"Authorization": self._password}, + json=data + ) as resp: if resp.status >= 300: raise NodeException(f"Getting errors from Lavalink REST api") @@ -274,11 +278,11 @@ class Node: return await resp.json() - def get_player(self, guild_id: int): + def get_player(self, guild_id: int) -> Optional[Player]: """Takes a guild ID as a parameter. Returns a voicelink Player object.""" return self._players.get(guild_id, None) - async def connect(self): + async def connect(self) -> Node: """Initiates a connection with a Lavalink node and adds it to the node pool.""" try: @@ -309,22 +313,22 @@ class Node: return self - async def disconnect(self): + async def disconnect(self) -> None: """Disconnects a connected Lavalink node and removes it from the node pool. This also destroys any players connected to the node. """ for player in self.players.copy().values(): await player.teardown() - if self._spotify_client_id and self._spotify_client_secret: - await self._spotify_client.close() + if self.spotify_client: + await self.spotify_client.close() await self._websocket.close() del self._pool._nodes[self._identifier] self._available = False self._task.cancel() - async def reconnect(self): + async def reconnect(self) -> None: await asyncio.sleep(10) for player in self.players.copy().values(): try: @@ -371,7 +375,7 @@ class Node: *, requester: Member, search_type: SearchType = SearchType.ytsearch - ): + ) -> Union[Track, Playlist]: """Fetches tracks from the node's REST api to parse into Lavalink. If you passed in Spotify API credentials, you can also pass in a @@ -385,15 +389,15 @@ class Node: query = f"{search_type}:{query}" if SPOTIFY_URL_REGEX.match(query): - if not self._spotify_client_id and not self._spotify_client_secret: - raise InvalidSpotifyClientAuthorization( + try: + if not self.spotify_client: + raise InvalidSpotifyClientAuthorization( "You did not provide proper Spotify client authorization credentials. " "If you would like to use the Spotify searching feature, " "please obtain Spotify API credentials here: https://developer.spotify.com/" ) - try: - spotify_results = await self._spotify_client.search(query=query) + spotify_results = await self.spotify_client.search(query=query) except Exception as _: raise TrackLoadError("Not able to find the provided Spotify entity, is it private?") @@ -507,7 +511,7 @@ class NodePool: This holds all the nodes that are to be used by the bot. """ - _nodes = {} + _nodes: Dict[str, Node] = {} def __repr__(self): return f"" @@ -518,7 +522,7 @@ class NodePool: return self._nodes @property - def node_count(self): + def node_count(self) -> Optional[Node]: return len(self._nodes.values()) @classmethod @@ -583,7 +587,6 @@ class NodePool: spotify_client_secret: Optional[str] = None, session: Optional[aiohttp.ClientSession] = None, resume_key: Optional[str] = None, - ) -> Node: """Creates a Node object to be then added into the node pool. For Spotify searching capabilites, pass in valid Spotify API credentials. diff --git a/voicelink/spotify/objects.py b/voicelink/spotify/objects.py index 2cdd5f9..2d55782 100644 --- a/voicelink/spotify/objects.py +++ b/voicelink/spotify/objects.py @@ -11,7 +11,7 @@ class Track: "uri" ) - def __init__(self, data: dict, image=None) -> None: + def __init__(self, data: dict, image = None) -> None: self.name: str = data.get('name', 'Unknown') self.artists: str = ", ".join(artist["name"] for artist in data.get('artists')) self.artist_id: list[str] = [artist['id'] for artist in data.get('artists')]