From 7ef2700730b807f8075471632156ba808d55ad93 Mon Sep 17 00:00:00 2001 From: Choco <94597336+ChocoMeow@users.noreply.github.com> Date: Wed, 26 Jun 2024 07:42:37 +0800 Subject: [PATCH] Optimized some code --- cogs/basic.py | 72 +++++++++++++++---------------------- cogs/playlist.py | 8 ++--- function.py | 6 ++-- ipc/methods.py | 61 ++++++++++++------------------- views/controller.py | 11 +++--- voicelink/enums.py | 11 ++++++ voicelink/player.py | 61 ++++++++++++++----------------- voicelink/pool.py | 35 ++++++++---------- voicelink/spotify/client.py | 2 +- voicelink/utils.py | 64 +++++++++++++++++++++++++-------- 10 files changed, 165 insertions(+), 166 deletions(-) diff --git a/cogs/basic.py b/cogs/basic.py index f985486..6e0f50b 100644 --- a/cogs/basic.py +++ b/cogs/basic.py @@ -31,6 +31,7 @@ from function import ( send, time as ctime, formatTime, + format_time, get_source, get_user, get_lang, @@ -329,15 +330,12 @@ class Basic(commands.Cog): if not player.is_privileged(ctx.author): if ctx.author in player.pause_votes: return await send(ctx, "voted", ephemeral=True) - else: - player.pause_votes.add(ctx.author) - if len(player.pause_votes) >= (required := player.required()): - pass - else: - return await send(ctx, "pauseVote", ctx.author, len(player.pause_votes), required) + + player.pause_votes.add(ctx.author) + if len(player.pause_votes) < (required := player.required()): + return await send(ctx, "pauseVote", ctx.author, len(player.pause_votes), required) await player.set_pause(True, ctx.author) - player.pause_votes.clear() await send(ctx, "paused", ctx.author) @commands.hybrid_command(name="resume", aliases=get_aliases("resume")) @@ -354,15 +352,12 @@ class Basic(commands.Cog): if not player.is_privileged(ctx.author): if ctx.author in player.resume_votes: return await send(ctx, "voted", ephemeral=True) - else: - player.resume_votes.add(ctx.author) - if len(player.resume_votes) >= (required := player.required()): - pass - else: - return await send(ctx, "resumeVote", ctx.author, len(player.resume_votes), required) + + player.resume_votes.add(ctx.author) + if len(player.resume_votes) < (required := player.required()): + return await send(ctx, "resumeVote", ctx.author, len(player.resume_votes), required) await player.set_pause(False, ctx.author) - player.resume_votes.clear() await send(ctx, "resumed", ctx.author) @commands.hybrid_command(name="skip", aliases=get_aliases("skip")) @@ -374,6 +369,9 @@ class Basic(commands.Cog): if not player: return await send(ctx, "noPlayer", ephemeral=True) + if not player.node._available: + return await send(ctx, "nodeReconnect") + if not player.is_playing: return await send(ctx, "skipError", ephemeral=True) @@ -384,19 +382,13 @@ class Basic(commands.Cog): return await send(ctx, "voted", ephemeral=True) else: player.skip_votes.add(ctx.author) - if len(player.skip_votes) >= (required := player.required()): - pass - else: + if len(player.skip_votes) < (required := player.required()): return await send(ctx, "skipVote", ctx.author, len(player.skip_votes), required) - if not player.node._available: - return await send(ctx, "nodeReconnect") - if index: player.queue.skipto(index) await send(ctx, "skipped", ctx.author) - if player.queue._repeat.mode == voicelink.LoopType.track: await player.set_repeat(voicelink.LoopType.off.name) @@ -411,18 +403,16 @@ class Basic(commands.Cog): if not player: return await send(ctx, "noPlayer", ephemeral=True) + if not player.node._available: + return await send(ctx, "nnodeReconnectode") + if not player.is_privileged(ctx.author): if ctx.author in player.previous_votes: return await send(ctx, "voted", ephemeral=True) - else: - player.previous_votes.add(ctx.author) - if len(player.previous_votes) >= (required := player.required()): - pass - else: - return await send(ctx, "backVote", ctx.author, len(player.previous_votes), required) - - if not player.node._available: - return await send(ctx, "nnodeReconnectode") + + player.previous_votes.add(ctx.author) + if len(player.previous_votes) < (required := player.required()): + return await send(ctx, "backVote", ctx.author, len(player.previous_votes), required) if not player.is_playing: player.queue.backto(index) @@ -432,7 +422,6 @@ class Basic(commands.Cog): await player.stop() await send(ctx, "backed", ctx.author) - if player.queue._repeat.mode == voicelink.LoopType.track: await player.set_repeat(voicelink.LoopType.off.name) @@ -451,8 +440,7 @@ class Basic(commands.Cog): if not player.current or player.position == 0: return await send(ctx, "noTrackPlaying", ephemeral=True) - num = formatTime(position) - if num is None: + if not (num := format_time(position)): return await send(ctx, "timeFormatError", ephemeral=True) await player.seek(num, ctx.author) @@ -679,9 +667,8 @@ class Basic(commands.Cog): if not player.current: return await send(ctx, "noTrackPlaying", ephemeral=True) - - num = formatTime(position) - if num is None: + + if not (num := format_time(position)): return await send(ctx, "timeFormatError", ephemeral=True) await player.seek(int(player.position + num)) @@ -702,8 +689,7 @@ class Basic(commands.Cog): if not player.current: return await send(ctx, "noTrackPlaying", ephemeral=True) - num = formatTime(position) - if num is None: + if not (num := format_time(position)): return await send(ctx, "timeFormatError", ephemeral=True) await player.seek(int(player.position - num)) @@ -737,12 +723,10 @@ class Basic(commands.Cog): if not player.is_privileged(ctx.author): if ctx.author in player.shuffle_votes: return await send(ctx, "voted", ephemeral=True) - else: - player.shuffle_votes.add(ctx.author) - if len(player.shuffle_votes) >= (required := player.required()): - pass - else: - return await send(ctx, "shuffleVote", ctx.author, len(player.shuffle_votes), required) + + player.shuffle_votes.add(ctx.author) + if len(player.shuffle_votes) < (required := player.required()): + return await send(ctx, "shuffleVote", ctx.author, len(player.shuffle_votes), required) await player.shuffle("queue", ctx.author) await send(ctx, "shuffled") diff --git a/cogs/playlist.py b/cogs/playlist.py index cd85282..47e66d2 100644 --- a/cogs/playlist.py +++ b/cogs/playlist.py @@ -21,7 +21,7 @@ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. """ -import discord, voicelink +import discord, voicelink, time from io import StringIO from discord import app_commands @@ -39,7 +39,6 @@ from function import ( logger ) -from datetime import datetime from views import PlaylistView, InboxView, HelpView def assign_playlist_id(existed: list) -> str: @@ -298,7 +297,7 @@ class Playlists(commands.Cog, name="playlist"): {"$push": {"inbox": { 'sender': ctx.author.id, 'referId': result['id'], - 'time': datetime.now(), + 'time': time.time(), 'title': f'Playlist invitation from {ctx.author}', 'description': f"You are invited to use this playlist.\nPlaylist Name: {result['playlist']['name']}\nPlaylist type: {result['playlist']['type']}", 'type': 'invite' @@ -357,7 +356,8 @@ class Playlists(commands.Cog, name="playlist"): await update_user(data['sender'], {"$push": {f"playlist.{data['referId']}.perms.read": ctx.author.id}}) update_data[f'playlist.{addId}'] = { 'user': data['sender'], 'referId': data['referId'], - 'name': f"Share{data['time'].strftime('%M%S')}", 'type': 'share' + 'name': f"Share{time.strftime('%M%S', time.gmtime(int(data['time'])))}", + 'type': 'share' } update_data["inbox"] = view.inbox dId.add(addId) diff --git a/function.py b/function.py index 458daf2..9dbf1b0 100644 --- a/function.py +++ b/function.py @@ -83,11 +83,11 @@ def time(millis:int) -> str: minutes=(millis/(1000*60))%60 hours=(millis/(1000*60*60))%24 if hours > 1: - return "%02d:%02d:%02d" % (hours, minutes, seconds) + return "%d:%02d:%02d" % (hours, minutes, seconds) else: return "%02d:%02d" % (minutes, seconds) -def formatTime(number:str) -> Optional[int]: +def format_time(number:str) -> int: try: try: num = strptime(number, '%M:%S') @@ -97,7 +97,7 @@ def formatTime(number:str) -> Optional[int]: except ValueError: num = strptime(number, '%H:%M:%S') except: - return None + return 0 return (int(num.tm_hour) * 3600 + int(num.tm_min) * 60 + int(num.tm_sec)) * 1000 diff --git a/ipc/methods.py b/ipc/methods.py index 0a289f6..e87b84c 100644 --- a/ipc/methods.py +++ b/ipc/methods.py @@ -85,11 +85,6 @@ async def initBot(bot: commands.Bot, data: Dict) -> Dict: async def initUser(bot: commands.Bot, data: Dict) -> Dict: user_id = int(data.get("user_id")) data = await func.get_user(user_id) - - for inbox in data.get("inbox", []): - dt = inbox["time"] - utc_time = dt.replace(tzinfo=timezone.utc) - inbox['time'] = utc_time.timestamp() return { "op": "initUser", @@ -150,11 +145,10 @@ async def skipTo(player: Player, member: Member, data: Dict) -> None: elif member in player.skip_votes: return error_msg(player.get_msg('voted'), user_id=member.id) + else: player.skip_votes.add(member) - if len(player.skip_votes) >= (required := player.required()): - pass - else: + if len(player.skip_votes) < (required := player.required()): return error_msg(player.get_msg('skipVote').format(member, len(player.skip_votes), required), guild_id=player.guild.id) index = data.get("index", 1) @@ -172,11 +166,10 @@ async def backTo(player: Player, member: Member, data: Dict) -> None: elif member in player.skip_votes: return error_msg(player.get_msg('voted'), user_id=member.id) + else: player.skip_votes.add(member) - if len(player.skip_votes) >= (required := player.required()): - pass - else: + if len(player.skip_votes) < (required := player.required()): return error_msg(player.get_msg('backVote').format(member, len(player.skip_votes), required), guild_id=player.guild.id) index = data.get("index", 1) @@ -220,32 +213,26 @@ async def addTracks(player: Player, member: Member, data: Dict) -> None: if not player.is_playing: await player.do_next() -async def getTracks(player: Player, member: Member, data: Dict) -> Dict: +async def getTracks(bot: commands.Bot, data: Dict) -> Dict: query = data.get("query", None) if query: - payload = {"op": "getTracks", "user_id": str(member.id)} - tracks = await player.get_tracks(query, requester=member) + payload = {"op": "getTracks", "user_id": data.get("user_id"), "callback": data.get("callback")} + tracks = await NodePool.get_node().get_tracks(query=query, requester=None) if not tracks: return payload - - if isinstance(tracks, Playlist): - tracks = [ track for track in tracks.tracks[:50] ] - payload["tracks"] = [ track.track_id for track in tracks ] + payload["tracks"] = [ track.track_id for track in (tracks.tracks if isinstance(tracks, Playlist) else tracks ) ] return payload async def shuffleTrack(player: Player, member: Member, data: Dict) -> None: if not player.is_privileged(member): - if member in player.shuffle_votes: return error_msg(player.get_msg('voted'), user_id=member.id) - else: - player.shuffle_votes.add(member) - if len(player.shuffle_votes) >= (required := player.required()): - pass - else: - return error_msg(player.get_msg('shuffleVote').format(member, len(player.skip_votes), required), guild_id=player.guild.id) + + player.shuffle_votes.add(member) + if len(player.shuffle_votes) < (required := player.required()): + return error_msg(player.get_msg('shuffleVote').format(member, len(player.skip_votes), required), guild_id=player.guild.id) await player.shuffle(data.get("type", "queue"), member) @@ -276,23 +263,19 @@ async def updatePause(player: Player, member: Member, data: Dict) -> None: if pause: if member in player.pause_votes: return error_msg(player.get_msg('voted'), user_id=member.id) - else: - player.pause_votes.add(member) - if len(player.pause_votes) >= (required := player.required()): - pass - else: - return error_msg(player.get_msg('pauseVote').format(member, len(player.pause_votes), required), guild_id=player.guild.id) + + player.pause_votes.add(member) + if len(player.pause_votes) < (required := player.required()): + return error_msg(player.get_msg('pauseVote').format(member, len(player.pause_votes), required), guild_id=player.guild.id) + else: if member in player.resume_votes: return error_msg(player.get_msg('voted'), user_id=member.id) - else: - player.resume_votes.add(member) - if len(player.resume_votes) >= (required := player.required()): - pass - else: - return error_msg(player.get_msg('resumeVote').format(member, len(player.resume_votes), required), guild_id=player.guild.id) + + player.resume_votes.add(member) + if len(player.resume_votes) < (required := player.required()): + return error_msg(player.get_msg('resumeVote').format(member, len(player.resume_votes), required), guild_id=player.guild.id) - player.pause_votes.clear() if pause else player.resume_votes.clear() await player.set_pause(pause, member) async def updatePosition(player: Player, member: Member, data: Dict) -> None: @@ -641,7 +624,7 @@ async def process_methods(ipc_client, bot: commands.Bot, data: Dict) -> None: RATELIMIT_COUNTER[user_id] = {"time": time.time(), "count": 0} else: - if RATELIMIT_COUNTER[user_id]["count"] >= 200: + if RATELIMIT_COUNTER[user_id]["count"] >= 100: return await ipc_client.send({"op": "rateLimited", "user_id": str(user_id)}) RATELIMIT_COUNTER[user_id]["count"] += method.credit diff --git a/views/controller.py b/views/controller.py index 678b9a4..8b4dbab 100644 --- a/views/controller.py +++ b/views/controller.py @@ -109,7 +109,6 @@ class Resume(ControlButton): if len(votes) < (required := self.player.required()): return await self.send(interaction, f"{vote_type}Vote", interaction.user, len(votes), required) - votes.clear() self.emoji = emoji if not self.disable_button_text: self.label = await func.get_lang(interaction.guild.id, button) @@ -391,7 +390,7 @@ class Tracks(discord.ui.Select): if self.player.settings.get("controller_msg", True): await func.send(interaction, "skipped", interaction.user) -btnType = { +BUTTONTYPE: Dict[str, ControlButton] = { "back": Back, "resume": Resume, "skip": Skip, @@ -408,7 +407,7 @@ btnType = { "rewind": Rewind } -btnColor = { +BUTTONCOLOR: Dict[str, discord.ButtonStyle] = { "blue": discord.ButtonStyle.primary, "grey": discord.ButtonStyle.secondary, "red": discord.ButtonStyle.danger, @@ -426,8 +425,8 @@ class InteractiveController(discord.ui.View): if isinstance(btn, Dict): color = list(btn.values())[0] btn = list(btn.keys())[0] - btnClass = btnType.get(btn.lower()) - style = btnColor.get(color.lower(), btnColor["grey"]) + btnClass = BUTTONTYPE.get(btn.lower()) + style = BUTTONCOLOR.get(color.lower(), BUTTONCOLOR["grey"]) if not btnClass or (self.player.queue.is_empty and btn == "tracks"): continue self.add_item(btnClass(player=player, style=style, row=row)) @@ -459,4 +458,4 @@ class InteractiveController(discord.ui.View): elif isinstance(error, Exception): await interaction.response.send_message(error) - return + return \ No newline at end of file diff --git a/voicelink/enums.py b/voicelink/enums.py index f3ada45..c8c53fa 100644 --- a/voicelink/enums.py +++ b/voicelink/enums.py @@ -59,6 +59,17 @@ class SearchType(Enum): def __str__(self) -> str: return self.value +class RequestMethod(Enum): + """The enum for the different request methods in Voicelink + """ + get = "get" + patch = "patch" + delete = "delete" + post = "post" + + def __str__(self) -> str: + return self.value + class NodeAlgorithm(Enum): """The enum for the different node algorithms in Voicelink. diff --git a/voicelink/player.py b/voicelink/player.py index 3dc1c2d..29ae273 100644 --- a/voicelink/player.py +++ b/voicelink/player.py @@ -42,7 +42,7 @@ from discord import ( from discord.ext import commands from . import events -from .enums import SearchType, LoopType +from .enums import SearchType, LoopType, RequestMethod from .events import VoicelinkEvent, TrackEndEvent, TrackStartEvent from .exceptions import VoicelinkException, FilterInvalidArgument, TrackInvalidPosition, TrackLoadError, FilterTagAlreadyInUse, DuplicateTrack from .filters import Filter, Filters @@ -255,6 +255,10 @@ class Player(VoiceProtocol): return manage_perm or (self.settings['dj'] in [role.id for role in user.roles]) return self.dj.id == user.id or manage_perm + async def send(self, method: RequestMethod, query: str = None, data: Union[Dict, str] = {}) -> Dict: + uri: str = f"sessions/{self._node._session_id}/players/{self._guild.id}" + (f"?{query}" if query else "") + return await self._node.send(method, query=uri, data=data) + async def _update_state(self, data: dict) -> None: state: dict = data.get("state") self._last_update = time.time() * 1000 @@ -284,11 +288,7 @@ class Player(VoiceProtocol): "sessionId": state['sessionId'], } - await self._node.send( - method=0, guild_id=self._guild.id, - data = {"voice": data} - ) - + await self.send(method=RequestMethod.patch, data={"voice": data}) self._logger.debug(f"Player in {self.guild.name}({self.guild.id}) dispatched voice update to {state['event']['endpoint']} with data {data}") async def on_voice_server_update(self, data: dict): @@ -471,7 +471,7 @@ class Player(VoiceProtocol): async def stop(self): """Stops the currently playing track.""" self._current = None - await self._node.send(method=0, guild_id=self._guild.id, data={'encodedTrack': None}) + await self.send(method=RequestMethod.patch, data={'encodedTrack': None}) async def disconnect(self, *, force: bool = False): """Disconnects the player from voice.""" @@ -495,7 +495,7 @@ class Player(VoiceProtocol): assert self.channel is None and not self.is_connected self._node._players.pop(self.guild.id) - await self._node.send(method=1, guild_id=self._guild.id) + await self.send(method=RequestMethod.delete) async def play( self, @@ -523,17 +523,10 @@ class Player(VoiceProtocol): data = { "encodedTrack": track.original.track_id if track.original else track.track_id, - "position": str(start) + "position": str(start if start else track.position), + "endTime": str(end if end else track.original.length) } - - if end > 0: - data["endTime"] = str(end) - - await self._node.send( - method=0, guild_id=self._guild.id, - data=data, - query=f"noReplace={ignore_if_playing}" - ) + await self.send(method=RequestMethod.patch, query=f"noReplace={ignore_if_playing}", data=data) self._current = track @@ -545,8 +538,7 @@ class Player(VoiceProtocol): async def add_track(self, raw_tracks: Union[Track, List[Track]], *, at_font: bool = False, duplicate: bool = True) -> int: tracks: List[Track] = [] - - _duplicate_tracks = () if self.queue._allow_duplicate and duplicate else (track.uri for track in self.queue._queue) + _duplicate_tracks = [] if self.queue._allow_duplicate and duplicate else [track.uri for track in self.queue._queue] raw_tracks = raw_tracks[0] if isinstance(raw_tracks, List) and len(raw_tracks) == 1 else raw_tracks try: @@ -554,8 +546,9 @@ class Player(VoiceProtocol): for track in raw_tracks: if track.uri in _duplicate_tracks: continue - self.queue.put_at_front(track) if at_font else self.queue.put(track) + self.queue.put_at_front(track) if at_front else self.queue.put(track) tracks.append(track) + _duplicate_tracks.append(track.uri) else: if raw_tracks.uri in _duplicate_tracks: raise DuplicateTrack(self.get_msg("voicelinkDuplicateTrack")) @@ -587,7 +580,7 @@ class Player(VoiceProtocol): if position < 0 or position > self._current.original.length: raise TrackInvalidPosition("Seek position must be between 0 and the track length") - await self._node.send(method=0, guild_id=self._guild.id, data={"position": position}) + await self.send(method=RequestMethod.patch, data={"position": position}) if self.is_ipc_connected: await self.send_ws({"op": "updatePosition", "position": position}, requester) @@ -596,17 +589,20 @@ class Player(VoiceProtocol): async def set_pause(self, pause: bool, requester: Member = None) -> bool: """Sets the pause state of the currently playing track.""" - await self._node.send(method=0, guild_id=self._guild.id, data={"paused": pause}) + + self._paused = pause + self.pause_votes.clear() if pause else self.resume_votes.clear() + await self.send(method=RequestMethod.patch, data={"paused": pause}) + if self.is_ipc_connected: await self.send_ws({"op": "updatePause", "pause": pause}, requester) - self._paused = pause - + self._logger.debug(f"Player in {self.guild.name}({self.guild.id}) has been {'paused' if pause else 'resumed'}.") return self._paused async def set_volume(self, volume: int, requester: Member = None) -> int: """Sets the volume of the player as an integer. Lavalink accepts values from 0 to 500.""" - await self._node.send(method=0, guild_id=self._guild.id, data={"volume": volume}) + await self.send(method=RequestMethod.patch, data={"volume": volume}) self._volume = volume if self.is_ipc_connected: @@ -626,11 +622,8 @@ class Player(VoiceProtocol): if self.is_ipc_connected: await self.send_ws({ "op": "shuffleTrack", - "tracks": [track.track_id for track in self.queue._queue], - "verified": { - "index": self.queue._position if queue_type == "queue" else 0, - "track_id": replacement[0].track_id, - } + "tracks": [{"track_id": track.track_id, "requester_id": str(track.requester.id)} for track in replacement], + "queue_type": queue_type }, requester) self._logger.debug(f"Player in {self.guild.name}({self.guild.id}) has been shuffled the queue.") @@ -679,7 +672,7 @@ class Player(VoiceProtocol): except FilterTagAlreadyInUse: raise FilterTagAlreadyInUse(self.get_msg("FilterTagAlreadyInUse")) payload = self._filters.get_all_payloads() - await self._node.send(method=0, guild_id=self._guild.id, data={"filters": payload}) + await self.send(method=RequestMethod.patch, data={"filters": payload}) if fast_apply: await self.seek(self.position) @@ -689,7 +682,7 @@ class Player(VoiceProtocol): async def remove_filter(self, filter_tag: str, fast_apply=False) -> Filters: self._filters.remove_filter(filter_tag=filter_tag) payload = self._filters.get_all_payloads() - await self._node.send(method=0, guild_id=self._guild.id, data={"filters": payload}) + await self.send(method=RequestMethod.patch, data={"filters": payload}) if fast_apply: await self.seek(self.position) @@ -700,7 +693,7 @@ class Player(VoiceProtocol): if not self._filters: raise FilterInvalidArgument("You must have filters applied first in order to use this method.") self._filters.reset_filters() - await self._node.send(method=0, guild_id=self._guild.id, data={"filters": {}}) + await self.send(method=RequestMethod.patch, data={"filters": {}}) if fast_apply: await self.seek(self.position) diff --git a/voicelink/pool.py b/voicelink/pool.py index 1b9fe0b..f29f314 100644 --- a/voicelink/pool.py +++ b/voicelink/pool.py @@ -50,7 +50,8 @@ from .exceptions import ( TrackLoadError ) from .objects import Playlist, Track -from .utils import ExponentialBackoff, NodeStats, Ping +from .utils import ExponentialBackoff, NodeStats, NodeInfo, Ping +from .enums import RequestMethod if TYPE_CHECKING: from .player import Player @@ -69,7 +70,6 @@ URL_REGEX = re.compile( ) NODE_VERSION = "v4" -CALL_METHOD = ["PATCH", "DELETE"] class Node: """The base class for a node. @@ -123,6 +123,7 @@ class Node: } self._players: Dict[int, Player] = {} + self._info: Optional[NodeInfo] = None self._spotify_client_id: Optional[str] = spotify_client_id self._spotify_client_secret: Optional[str] = spotify_client_secret @@ -136,6 +137,10 @@ class Node: f"player_count={len(self._players)}>" ) + 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) + @property def spotify_client(self) -> Optional[spotify.Client]: if not self._spotify_client: @@ -185,7 +190,7 @@ class Node: return Ping(self._host, port=self._port).get_ping() async def _update_handler(self, data: dict) -> None: - #await self._bot.wait_until_ready() + await self._bot.wait_until_ready() if not data: return @@ -252,22 +257,13 @@ 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] = {} - ) -> dict: + async def send(self, method: RequestMethod, query: str, data: Union[dict, str] = {}) -> dict: if not self._available: 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}/{query}" async with self._session.request( - method=CALL_METHOD[method], + method=method.value, url=uri, headers={"Authorization": self._password}, json=data @@ -275,15 +271,11 @@ class Node: if resp.status >= 300: raise NodeException(f"Getting errors from Lavalink REST api") - if method == CALL_METHOD[1]: + if method == RequestMethod.delete: return await resp.json(content_type=None) return await resp.json() - 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) -> Node: """Initiates a connection with a Lavalink node and adds it to the node pool.""" @@ -294,7 +286,8 @@ class Node: self._task = self._bot.loop.create_task(self._listen()) self._available = True - + self._info = NodeInfo(await self.send(RequestMethod.get, query="info")) + self._logger.info(f"Node [{self._identifier}] is connected!") except aiohttp.ClientConnectorError: diff --git a/voicelink/spotify/client.py b/voicelink/spotify/client.py index 186ee76..209e36a 100644 --- a/voicelink/spotify/client.py +++ b/voicelink/spotify/client.py @@ -146,7 +146,7 @@ class Client: async def get_categories(self) -> List[Category]: if not self._categories: - request_url = BASE_URL + "browse/categories" + request_url = f"{BASE_URL}browse/categories" data = await self.get_request(request_url) self._categories = [Category(item) for item in data.get("items", [])] diff --git a/voicelink/utils.py b/voicelink/utils.py index 1b0025e..3a6c4bc 100644 --- a/voicelink/utils.py +++ b/voicelink/utils.py @@ -27,9 +27,15 @@ import socket from timeit import default_timer as timer from itertools import zip_longest +from typing import Dict, Optional + __all__ = [ "ExponentialBackoff", - "NodeStats" + "NodeStats", + "NodeInfoVersion", + "NodeInfo", + "Plugin", + "Ping" ] class ExponentialBackoff: @@ -85,26 +91,56 @@ class NodeStats: Gives critical information on the node, which is updated every minute. """ - def __init__(self, data: dict) -> None: + def __init__(self, data: Dict) -> None: - memory: dict = data.get("memory") - self.used = memory.get("used") - self.free = memory.get("free") - self.reservable = memory.get("reservable") - self.allocated = memory.get("allocated") + memory: Dict = data.get("memory") + self.used: int = memory.get("used") + self.free: int = memory.get("free") + self.reservable: int = memory.get("reservable") + self.allocated: int = memory.get("allocated") - cpu: dict = data.get("cpu") - self.cpu_cores = cpu.get("cores") - self.cpu_system_load = cpu.get("systemLoad") - self.cpu_process_load = cpu.get("lavalinkLoad") + cpu: Dict = data.get("cpu") + self.cpu_cores: int = cpu.get("cores") + self.cpu_system_load: float = cpu.get("systemLoad") + self.cpu_process_load: float = cpu.get("lavalinkLoad") - self.players_active = data.get("playingPlayers") - self.players_total = data.get("players") - self.uptime = data.get("uptime") + self.players_active: int = data.get("playingPlayers") + self.players_total: int = data.get("players") + self.uptime: int = data.get("uptime") def __repr__(self) -> str: return f"" +class NodeInfoVersion: + """The base class for the node info object. + Gives version information on the node. + """ + def __init__(self, data: Dict) -> None: + self.semver: str = data.get("semver") + self.major: int = data.get("major") + self.minor: int = data.get("minor") + self.patch: int = data.get("patch") + self.pre_release: Optional[str] = data.get("preRelease") + self.build: Optional[str] = data.get("build") + +class NodeInfo: + """The base class for the node info object. + Gives basic information on the node. + """ + def __init__(self, data: Dict) -> None: + self.version: NodeInfoVersion = NodeInfoVersion(data.get("version")) + self.build_time: int = data.get("buildTime") + self.jvm: str = data.get("jvm") + self.lavaplayer: str = data.get("lavaplayer") + self.plugins: Optional[Dict[str, Plugin]] = [Plugin(plugin_data) for plugin_data in data.get("plugins")] + +class Plugin: + """The base class for the plugin object. + Gives basic information on the plugin. + """ + def __init__(self, data: Dict) -> None: + self.name: str = data.get("name") + self.version: str = data.get("version") class Ping: # Thanks to https://github.com/zhengxiaowai/tcping for the nice ping impl