From ebb07960bf881d20e20b4c4a69b648f283fa64b0 Mon Sep 17 00:00:00 2001 From: strNophix Date: Fri, 19 Nov 2021 21:10:29 +0100 Subject: [PATCH 1/5] database rewrite --- bot.py | 24 +++++ cogs/information.py | 9 +- cogs/music.py | 14 ++- cogs/settings.py | 148 +++++++++++++++++++++++++++---- context.py | 37 ++++++++ scripts/channel_to_redis.py | 6 +- scripts/playlist_to_redis.py | 7 +- tunebot/__init__.py | 1 + tunebot/abc.py | 58 ++++++++++++ tunebot/redis/__init__.py | 4 + tunebot/redis/autojoin.py | 41 +++++++++ tunebot/redis/entity.py | 49 ++++++++++ tunebot/redis/playlist.py | 33 +++++++ tunebot/redis/playlist_source.py | 35 ++++++++ utils/EmbedGenerator.py | 5 +- utils/classes.py | 2 + utils/database.py | 35 -------- utils/exceptions.py | 3 +- 18 files changed, 435 insertions(+), 76 deletions(-) create mode 100644 tunebot/__init__.py create mode 100644 tunebot/abc.py create mode 100644 tunebot/redis/__init__.py create mode 100644 tunebot/redis/autojoin.py create mode 100644 tunebot/redis/entity.py create mode 100644 tunebot/redis/playlist.py create mode 100644 tunebot/redis/playlist_source.py delete mode 100644 utils/database.py diff --git a/bot.py b/bot.py index 3af0a2c..9b3389c 100644 --- a/bot.py +++ b/bot.py @@ -4,6 +4,7 @@ from typing import Any from typing import Dict from typing import List from typing import Sequence +from typing import TYPE_CHECKING import aioredis import discord @@ -21,6 +22,14 @@ from discord.ext.commands.errors import ExtensionNotFound from discord.ext.commands.errors import NoEntryPointError from context import CustomContext +from tunebot.redis import GlobalRedisAutoJoin +from tunebot.redis import GlobalRedisPlaylist +from tunebot.redis import GlobalRedisPlaylistSource + +if TYPE_CHECKING: + from tunebot import GlobalPlaylist + from tunebot import GlobalPlaylistSource + from tunebot import GlobalAutoJoin class TuneBot(commands.Bot): @@ -39,9 +48,24 @@ class TuneBot(commands.Bot): self.initial_cog_names: List[str] = self.config.get("cogs", []) self.colors: Dict[str, Color] = self.process_colours(config.get("colors", [])) + self.redis_prefix = self.config["redis_prefix"] self._redis_client: Redis = aioredis.from_url( self.config["redis_url"], encoding="utf-8", decode_responses=True ) + + self.global_autojoin: GlobalAutoJoin = GlobalRedisAutoJoin( + self._redis_client, + self.redis_prefix, + ) + self.global_playlist: GlobalPlaylist = GlobalRedisPlaylist( + self._redis_client, + self.redis_prefix, + ) + self.global_playlist_source: GlobalPlaylistSource = GlobalRedisPlaylistSource( + self._redis_client, + self.redis_prefix, + ) + self.invite_link: str = "" slash_guilds = None diff --git a/cogs/information.py b/cogs/information.py index bbdab08..f20f2be 100644 --- a/cogs/information.py +++ b/cogs/information.py @@ -6,14 +6,12 @@ import discord import humanize import lavalink from discord.ext import commands -from discord.ext import tasks from discord.ext.commands import Context from bot import TuneBot from context import CustomContext -from utils.database import AutoJoin -from utils.EmbedGenerator import EmbedGenerator from utils.classes import BaseCog +from utils.EmbedGenerator import EmbedGenerator from utils.paginator import HelpPaginator @@ -63,16 +61,15 @@ class InformationCog(BaseCog, name="Information"): fmt = ( f"**Lavalink:** `{lavalink.__version__}`\n\n" f"Connected to `{len(self.bot.lavalink.node_manager.available_nodes)}` nodes.\n" - #f"Best available Node `{self.bot.lavalink.node_manager.find_ideal_node().name.__repr__()}`\n" + # f"Best available Node `{self.bot.lavalink.node_manager.find_ideal_node().name.__repr__()}`\n" f"`{len(self.bot.lavalink.player_manager.players)}` players are distributed on nodes.\n" f"`{sum([n.stats.players for n in nodes])}` players are distributed on server.\n" f"`{sum([n.stats.playing_players for n in nodes])}` players are playing on server.\n\n" f"Server Memory: `{used}/{total}` | `({free} free)`\n" f"Server CPU: `{cpu}`\n\n" - #f"Server Uptime: `{datetime.timedelta(milliseconds=node.stats.uptime)}`" + # f"Server Uptime: `{datetime.timedelta(milliseconds=node.stats.uptime)}`" ) await ctx.send(fmt) - AutoJoin.get_channels() @commands.command(name="help", aliases=["about", "info"]) @commands.cooldown(1, 1, commands.BucketType.user) diff --git a/cogs/music.py b/cogs/music.py index 02cb4ee..b93d50d 100644 --- a/cogs/music.py +++ b/cogs/music.py @@ -4,24 +4,20 @@ import re from typing import Optional import discord -from discord.channel import TextChannel -from discord.ext.commands.context import Context -from discord.ext.commands.errors import CommandError import lavalink from discord import Embed from discord.channel import TextChannel from discord.ext import commands from discord.ext.commands.context import Context +from discord.ext.commands.errors import CommandError from lavalink.models import AudioTrack from lavalink.models import DefaultPlayer from bot import TuneBot -from utils.classes import BaseCog from context import CustomContext -from utils.database import AutoJoin -from utils.database import Playlist -from utils.exceptions import EmbeddedCommandException +from utils.classes import BaseCog from utils.EmbedGenerator import EmbedGenerator +from utils.exceptions import EmbeddedCommandException url_rx = re.compile(r"https?://(?:www\.)?.+") @@ -92,7 +88,7 @@ class Music(BaseCog): await self.async_init() async def async_init(self): - redis_result = await AutoJoin.get_channels(self.bot._redis_client) + redis_result = await self.bot.global_autojoin.fetch_channels() while len(self.bot.lavalink.node_manager.available_nodes) == 0: await asyncio.sleep(1) @@ -112,7 +108,7 @@ class Music(BaseCog): await textchannel.send("Automatically joined the voice channel") async def fill_player_queue(self, player: DefaultPlayer, buffer: Optional[int] = 1): - queries = await Playlist.random(self.bot._redis_client, buffer) + queries = await self.bot.global_playlist.pick_random(buffer) # Get the results for the query from Lavalink. for query in queries: result = await player.node.get_tracks(query) diff --git a/cogs/settings.py b/cogs/settings.py index fee02a4..8f37ddf 100644 --- a/cogs/settings.py +++ b/cogs/settings.py @@ -1,13 +1,9 @@ +from discord import Message from discord.ext import commands -from utils.EmbedGenerator import EmbedGenerator -from utils.classes import BaseCog -from utils.database import AutoJoin -from bot import TuneBot -from discord.ext.commands import Context from bot import TuneBot from context import CustomContext -from utils.database import AutoJoin +from utils.classes import BaseCog from utils.EmbedGenerator import EmbedGenerator @@ -19,28 +15,146 @@ class SettingsCog(BaseCog, name="Settings"): await EmbedGenerator.Message( ctx, "Autojoin", - f"Usage:\n\n`{ctx.prefix}autojoin set`\n`{ctx.prefix}autojoin unset`", + f"Usage:\n\n`{ctx.prefix}autojoin enable`\n`{ctx.prefix}autojoin disable`", ) - @autojoin.command(name="enable") + @autojoin.command(name="enable", aliases=["set"]) @commands.has_permissions(manage_channels=True) @commands.cooldown(rate=1, per=5, type=commands.BucketType.user) async def autojoin_set(self, ctx: CustomContext): """Enable the bot automatically joining""" - voicechannel_id = ctx.author.voice.channel.id - textchannel_id = ctx.message.channel.id - await AutoJoin.update_channel( - ctx.redis, ctx.guild.id, voicechannel_id, textchannel_id - ) - await EmbedGenerator.Message(ctx, "Autojoin", "`enabled`") + voice_state = ctx.author.voice + if not voice_state: + embed = ctx.create_embed() + embed.title = "Please join a voice channel before running this command." + await ctx.send(embed=embed) + return - @autojoin.command(name="disable") + await ctx.autojoin.update(voice_state.channel.id, ctx.message.channel.id) + embed = ctx.create_embed() + embed.title = f"AutoJoin enabled for #{voice_state.channel.name}" + await ctx.send(embed=embed) + + @autojoin.command(name="disable", aliases=["unset"]) @commands.has_permissions(manage_channels=True) @commands.cooldown(rate=1, per=5, type=commands.BucketType.user) async def autojoin_del(self, ctx: CustomContext): """Disable the bot automatically joining""" - await AutoJoin.del_channel(ctx.redis, ctx.guild.id) - await EmbedGenerator.Message(ctx, "Autojoin", "`disabled`") + await ctx.autojoin.disable() + embed = ctx.create_embed() + embed.title = f"AutoJoin disabled" + await ctx.send(embed=embed) + + @commands.is_owner() + @commands.group( + name="source", + aliases=["src"], + invoke_without_command=True, + slash_command=False, + hidden=True, + ) + async def source(self, ctx: CustomContext): + """Displays all possible options for the `source` command""" + prefix = self.bot.config["prefixes"][0] + embed = ctx.create_embed() + embed.title = "All options:" + embed.description = f"```{prefix}source list\n{prefix}source add \n{prefix}source remove \n{prefix}source sync```" + await ctx.send(embed=embed) + + @commands.is_owner() + @source.command(name="remove") + async def source_remove(self, ctx: CustomContext, source_url: str): + """Removes a source from the bot""" + if await ctx.playlist_source.remove(source_url): + prefix = self.bot.config["prefixes"][0] + embed = ctx.create_embed() + embed.title = "Removed source succesfully" + embed.description = ( + f"Please use `{prefix}source sync` to persist these changes." + ) + await ctx.send(embed=embed) + return + + embed = ctx.create_embed() + embed.title = "Could not remove source, the specified source might not exist" + await ctx.send(embed=embed) + + @commands.is_owner() + @source.command(name="add") + async def source_add(self, ctx: CustomContext, source_url: str): + """ + Add's a source to the bot + + Supported sources: YouTube, SoundCloud, Bandcamp, Vimeo, Twitch and HTTP(S) URL's + """ + embed = ctx.create_embed() + embed.title = "Started processing source" + message = await ctx.send(embed=embed) + + query_result: Any = await self.bot.lavalink.get_tracks(source_url) + + if query_result["loadType"] == "LOAD_FAILED": + embed = ctx.create_embed() + embed.title = "The specified URL is not a valid source" + embed.description = f"Supported sources: YouTube, SoundCloud, Bandcamp, Vimeo, Twitch and HTTP(S) URL's" + if isinstance(message, Message): + await message.edit(embed=embed) + else: + await ctx.send(embed=embed) + return + + await ctx.playlist_source.add(source_url) + track_urls = [str(track["info"]["uri"]) for track in query_result["tracks"]] + await self.bot.global_playlist.add_tracks(track_urls) + + embed = ctx.create_embed() + embed.title = "Finished processing source" + embed.description = f"Added {len(track_urls)} tracks" + if isinstance(message, Message): + await message.edit(embed=embed) + else: + await ctx.send(embed=embed) + + @commands.is_owner() + @source.command(name="list", aliases=["ls"]) + async def source_list(self, ctx: CustomContext): + """Display a list of sources""" + # TODO: Implement pagination for sources + sources = await self.bot.global_playlist_source.fetch_sources() + if len(sources) > 0: + description = "\n".join([f"[{source}]({source})" for source in sources]) + else: + description = "This bot has no sources yet" + + embed = ctx.create_embed() + embed.title = "All sources:" + embed.description = description + await ctx.send(embed=embed) + + @commands.is_owner() + @source.command(name="sync") + async def source_sync(self, ctx: CustomContext): + """Forcefully resyncs all sources""" + failed_sources: list[str] = [] + await self.bot.global_playlist.clear() + sources = await self.bot.global_playlist_source.fetch_sources() + for source_url in sources: + query_result: Any = await self.bot.lavalink.get_tracks(source_url) + if query_result["loadType"] == "LOAD_FAILED": + failed_sources.append(source_url) + continue + + track_urls = [str(track["info"]["uri"]) for track in query_result["tracks"]] + await self.bot.global_playlist.add_tracks(track_urls) + + embed = ctx.create_embed() + embed.title = f"Finished sync ({len(failed_sources)} issues)" + if len(failed_sources) > 0: + embed.description = "\n".join( + [f"[{source_url}]({source_url})" for source_url in failed_sources] + ) + + await ctx.send(embed=embed) def setup(bot: TuneBot): diff --git a/context.py b/context.py index 3f41dc7..ca67efe 100644 --- a/context.py +++ b/context.py @@ -1,10 +1,23 @@ +from typing import TYPE_CHECKING + from aioredis.client import Redis +from discord import Embed from discord.ext import commands from discord.ext.commands.errors import CommandInvokeError from lavalink.models import DefaultPlayer +from tunebot.redis import RedisAutoJoin +from tunebot.redis import RedisPlaylistSource + +if TYPE_CHECKING: + from bot import TuneBot + from tunebot import PlaylistSource + from tunebot import AutoJoin + class CustomContext(commands.Context): + bot: "TuneBot" + @property def redis(self) -> Redis: return self.bot._redis_client @@ -15,3 +28,27 @@ class CustomContext(commands.Context): return self.bot.lavalink.player_manager.get(self.guild.id) raise CommandInvokeError("Lavalink is still starting up.") + + def create_embed(self) -> Embed: + bot: TuneBot = self.bot + color = bot.colors["embed"] + + avatar = None + if avatar_asset := self.author.avatar: + avatar = avatar_asset.with_static_format("jpeg") + + embed = Embed(color=color) + embed.set_footer(text=f"Requested by: {self.author}", icon_url=avatar) + return embed + + @property + def playlist_source(self) -> "PlaylistSource": + if not hasattr(self, "_playlist_source"): + self._playlist_source = RedisPlaylistSource(self) + return self._playlist_source + + @property + def autojoin(self) -> "AutoJoin": + if not hasattr(self, "_autojoin"): + self._autojoin = RedisAutoJoin(self) + return self._autojoin diff --git a/scripts/channel_to_redis.py b/scripts/channel_to_redis.py index 4f10e25..4c4315e 100644 --- a/scripts/channel_to_redis.py +++ b/scripts/channel_to_redis.py @@ -1,8 +1,9 @@ -import subprocess import json +import subprocess +import sys + import redis from yt_dlp import YoutubeDL -import sys if len(sys.argv) < 2: raise Exception("Expected youtube playlist/channel/video") @@ -25,4 +26,3 @@ for vid_url in vid_urls.split("\n"): vid_url = "https://www.youtube.com/watch?v=" + vid_url redis_client.sadd(f"{redis_prefix}:playlist", vid_url) print(vid_url) - diff --git a/scripts/playlist_to_redis.py b/scripts/playlist_to_redis.py index b5fecd7..be97a7a 100644 --- a/scripts/playlist_to_redis.py +++ b/scripts/playlist_to_redis.py @@ -1,8 +1,9 @@ import json -from aiotube import Playlist -import redis import sys +import redis +from aiotube import Playlist + if len(sys.argv) < 2: raise Exception("Expected path to file as argument") @@ -18,4 +19,4 @@ for line in file: for vid in playlist.videos(): yt_url = vid.url redis_client.sadd(f"{redis_prefix}:playlist", yt_url) - print(yt_url) \ No newline at end of file + print(yt_url) diff --git a/tunebot/__init__.py b/tunebot/__init__.py new file mode 100644 index 0000000..85439a2 --- /dev/null +++ b/tunebot/__init__.py @@ -0,0 +1 @@ +from tunebot.abc import * diff --git a/tunebot/abc.py b/tunebot/abc.py new file mode 100644 index 0000000..d8c580f --- /dev/null +++ b/tunebot/abc.py @@ -0,0 +1,58 @@ +from abc import ABC +from abc import abstractmethod +from typing import Optional + + +class GlobalPlaylistSource(ABC): + @abstractmethod + async def fetch_sources(self) -> set[str]: + pass + + +class PlaylistSource(ABC): + @abstractmethod + async def add(self, source_url: str): + pass + + @abstractmethod + async def remove(self, source_url: str) -> bool: + pass + + +class GlobalPlaylist(ABC): + @abstractmethod + async def pick_random(self, amount: Optional[int] = 1) -> set[str]: + pass + + @abstractmethod + async def add_tracks(self, urls: list[str]): + pass + + @abstractmethod + async def clear(self): + pass + + +class GlobalAutoJoin(ABC): + @abstractmethod + async def fetch_channels(self) -> dict[str, list[str]]: + pass + + +class AutoJoin(ABC): + @abstractmethod + async def update(self, voice_channel_id: int, text_channel_id: int): + pass + + @abstractmethod + async def disable(self): + pass + + +__all__ = ( + "GlobalPlaylistSource", + "PlaylistSource", + "GlobalPlaylist", + "GlobalAutoJoin", + "AutoJoin", +) diff --git a/tunebot/redis/__init__.py b/tunebot/redis/__init__.py new file mode 100644 index 0000000..5438850 --- /dev/null +++ b/tunebot/redis/__init__.py @@ -0,0 +1,4 @@ +from tunebot.redis.autojoin import * +from tunebot.redis.entity import * +from tunebot.redis.playlist import * +from tunebot.redis.playlist_source import * diff --git a/tunebot/redis/autojoin.py b/tunebot/redis/autojoin.py new file mode 100644 index 0000000..cbcb6c2 --- /dev/null +++ b/tunebot/redis/autojoin.py @@ -0,0 +1,41 @@ +from tunebot import AutoJoin +from tunebot import GlobalAutoJoin +from tunebot.redis import RedisBotEntity +from tunebot.redis import RedisContextEntity + + +class GlobalRedisAutoJoin(RedisBotEntity, GlobalAutoJoin): + async def fetch_channels(self) -> dict[str, list[str]]: + """ + Retrieves all guilds (with their configurations) where AutoJoin is enabled + """ + channels: dict[str, str] = await self.redis.hgetall(self.key("autojoin")) + return {key: value.split(":") for key, value in channels.items()} + + +class RedisAutoJoin(RedisContextEntity, AutoJoin): + async def update(self, voice_channel_id: int, text_channel_id: int): + """ + Upserts the configuration of an AutoJoin guild. + + Args: + voice_channel_id (int): [description] + text_channel_id (int): [description] + """ + if not self.ctx.guild: + raise Exception("This method can only be invoked inside of a guild.") + + value = f"{voice_channel_id}:{text_channel_id}" + await self.redis.hset(self.key("autojoin"), self.ctx.guild.id, value) + + async def disable(self): + """ + Removes the AutoJoin configuration of a guild. + """ + if not self.ctx.guild: + raise Exception("This method can only be invoked inside of a guild.") + + await self.redis.hdel(self.key("autojoin"), self.ctx.guild.id) + + +__all__ = ("GlobalRedisAutoJoin", "RedisAutoJoin") diff --git a/tunebot/redis/entity.py b/tunebot/redis/entity.py new file mode 100644 index 0000000..21f15b1 --- /dev/null +++ b/tunebot/redis/entity.py @@ -0,0 +1,49 @@ +from abc import ABC +from abc import abstractmethod +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from context import CustomContext + +from aioredis.client import Redis + + +class RedisEntity(ABC): + @property + @abstractmethod + def redis() -> Redis: + pass + + @abstractmethod + def key(self, name: str) -> str: + pass + + +class RedisBotEntity(RedisEntity): + def __init__(self, redis: Redis, prefix: str) -> None: + self._redis = redis + self.prefix = prefix + super().__init__() + + @property + def redis(self) -> Redis: + return self._redis + + def key(self, name: str) -> str: + return ":".join([self.prefix, name]) + + +class RedisContextEntity(RedisEntity): + def __init__(self, ctx: "CustomContext") -> None: + self.ctx = ctx + super().__init__() + + @property + def redis(self) -> Redis: + return self.ctx.redis + + def key(self, name: str) -> str: + return ":".join([self.ctx.bot.redis_prefix, name]) + + +__all__ = ("RedisEntity", "RedisBotEntity", "RedisContextEntity") diff --git a/tunebot/redis/playlist.py b/tunebot/redis/playlist.py new file mode 100644 index 0000000..bf90969 --- /dev/null +++ b/tunebot/redis/playlist.py @@ -0,0 +1,33 @@ +from typing import Optional + +from tunebot.abc import GlobalPlaylist +from tunebot.redis import RedisBotEntity + + +class GlobalRedisPlaylist(RedisBotEntity, GlobalPlaylist): + async def pick_random(self, amount: Optional[int] = 1) -> set[str]: + """ + Picks an amount of random tracks from the playlist in Redis + + Args: + amount (Optional[int], optional): [description]. Defaults to 1. + """ + return await self.redis.srandmember(self.key("playlist"), amount) + + async def add_tracks(self, urls: list[str]): + """ + Adds one or more urls to the playlist in Redis + + Args: + urls (list[str]): [description] + """ + await self.redis.sadd(self.key("playlist"), *urls) + + async def clear(self): + """ + Clears the entire playlist in Redis + """ + await self.redis.delete(self.key("playlist")) + + +__all__ = ("GlobalRedisPlaylist",) diff --git a/tunebot/redis/playlist_source.py b/tunebot/redis/playlist_source.py new file mode 100644 index 0000000..2fda47a --- /dev/null +++ b/tunebot/redis/playlist_source.py @@ -0,0 +1,35 @@ +from tunebot import GlobalPlaylistSource +from tunebot import PlaylistSource +from tunebot.redis import RedisBotEntity +from tunebot.redis import RedisContextEntity + + +class GlobalRedisPlaylistSource(RedisBotEntity, GlobalPlaylistSource): + async def fetch_sources(self) -> set[str]: + """ + Fetches all sources + """ + return await self.redis.smembers(self.key("sources")) + + +class RedisPlaylistSource(RedisContextEntity, PlaylistSource): + async def add(self, source_url: str): + """ + Adds a playlist source to Redis + + Args: + source_url (str): [description] + """ + await self.redis.sadd(self.key("sources"), source_url) + + async def remove(self, source_url: str) -> bool: + """ + Removes a playlist source from Redis + + Args: + source_url (str): [description] + """ + return await self.redis.srem(self.key("sources"), source_url) + + +__all__ = ("GlobalRedisPlaylistSource", "RedisPlaylistSource") diff --git a/utils/EmbedGenerator.py b/utils/EmbedGenerator.py index 95cf8c0..1ac67e5 100644 --- a/utils/EmbedGenerator.py +++ b/utils/EmbedGenerator.py @@ -1,8 +1,9 @@ -from typing import Optional, Union +from typing import Optional +from typing import Union import discord -from discord.ext.commands import Context from discord import Embed +from discord.ext.commands import Context class EmbedGenerator: diff --git a/utils/classes.py b/utils/classes.py index 5c90b02..9cc95b2 100644 --- a/utils/classes.py +++ b/utils/classes.py @@ -1,5 +1,7 @@ from typing import Dict + from discord.ext.commands import Cog + from bot import TuneBot diff --git a/utils/database.py b/utils/database.py deleted file mode 100644 index 0d48ec6..0000000 --- a/utils/database.py +++ /dev/null @@ -1,35 +0,0 @@ -from typing import Dict -from typing import List -from typing import Optional - -from aioredis import Redis - -from bot import redis_prefix - - -class AutoJoin: - @staticmethod - async def get_channels(redis: Redis) -> Dict[str, str]: - # return all channels - channels = await redis.hgetall(f"{redis_prefix}:autojoin") - return {key: value.split("-") for key, value in channels.items()} - - @staticmethod - async def update_channel( - redis: Redis, guild_id: int, voice_channel_id: int, text_channel_id: int - ): - await redis.hset( - f"{redis_prefix}:autojoin", - guild_id, - f"{voice_channel_id}-{text_channel_id}", - ) - - @staticmethod - async def del_channel(redis: Redis, guild_id: int): - await redis.hdel(f"{redis_prefix}:autojoin", guild_id) - - -class Playlist: - @staticmethod - async def random(redis: Redis, amount: Optional[int] = 1) -> List[str]: - return await redis.srandmember(f"{redis_prefix}:playlist", amount) diff --git a/utils/exceptions.py b/utils/exceptions.py index d807e17..b6a08e0 100644 --- a/utils/exceptions.py +++ b/utils/exceptions.py @@ -1,7 +1,8 @@ from discord.embeds import Embed -from context import CustomContext from discord.ext.commands import CommandError +from context import CustomContext + class EmbeddedCommandException(CommandError): def __init__(self, embed: Embed) -> None: From d2ecbbf6ae65b62adb0258e7395cad8ed367a26f Mon Sep 17 00:00:00 2001 From: strNophix Date: Fri, 19 Nov 2021 21:58:47 +0100 Subject: [PATCH 2/5] prevent circular import in tunebot.redis.__init__.py --- tunebot/redis/__init__.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tunebot/redis/__init__.py b/tunebot/redis/__init__.py index 5438850..1e4a501 100644 --- a/tunebot/redis/__init__.py +++ b/tunebot/redis/__init__.py @@ -1,4 +1,4 @@ +from tunebot.redis.entity import * # noreorder from tunebot.redis.autojoin import * -from tunebot.redis.entity import * from tunebot.redis.playlist import * from tunebot.redis.playlist_source import * From 8504d6abfdbdaec1a66a4e7f452d25a1a8ff17b5 Mon Sep 17 00:00:00 2001 From: strNophix Date: Sat, 20 Nov 2021 00:37:17 +0100 Subject: [PATCH 3/5] decoupled lavalink creation from cogs.music --- bot.py | 6 ++++++ cogs/music.py | 22 ++++------------------ context.py | 6 +++++- utils/classes.py | 7 +++++++ utils/database.py | 4 ++-- 5 files changed, 24 insertions(+), 21 deletions(-) diff --git a/bot.py b/bot.py index 3af0a2c..32cca39 100644 --- a/bot.py +++ b/bot.py @@ -87,6 +87,12 @@ class TuneBot(commands.Bot): print(f"Version: {discord.__version__}") print(f"Invite: {self.invite_link}") + ll = self.config["lavalink"] + self.lavalink = lavalink.Client(self.user.id) + self.lavalink.add_node( + ll["host"], ll["port"], ll["password"], ll["region"], ll["name"] + ) + def process_colours(self, colors: Dict[str, str]) -> Dict[str, Color]: colour_dict: Dict[str, Color] = {} for name, color in colors.items(): diff --git a/cogs/music.py b/cogs/music.py index 282cb67..436cce8 100644 --- a/cogs/music.py +++ b/cogs/music.py @@ -76,25 +76,11 @@ class LavalinkVoiceClient(discord.VoiceClient): class Music(BaseCog): @commands.Cog.listener() async def on_ready(self): - if not hasattr( - self.bot, "lavalink" - ): # This ensures the client isn't overwritten during cog reloads. - self.bot.lavalink = lavalink.Client(self.bot.user.id) - - ll = self.bot.config["lavalink"] - self.bot.lavalink.add_node( - ll["host"], ll["port"], ll["password"], ll["region"], ll["name"] - ) - - self.bot.lavalink.add_event_hook(self.track_hook) - await self.async_init() - - async def async_init(self): - redis_result = await AutoJoin.get_channels(self.bot._redis_client) - - while len(self.bot.lavalink.node_manager.available_nodes) == 0: + while not self.is_lavalink_ready(): await asyncio.sleep(1) + self.bot.lavalink.add_event_hook(self.track_hook) + redis_result = await AutoJoin.get_channels(self.bot._redis_client) for guild_id, (voicechannel_id, textchannel_id) in redis_result.items(): player = self.bot.lavalink.player_manager.create(guild_id) player.store("channel", textchannel_id) @@ -153,7 +139,7 @@ class Music(BaseCog): # This is essentially the same as `@commands.guild_only()` # except it saves us repeating ourselves (and also a few lines). - if not hasattr(self.bot, "lavalink"): + if not self.is_lavalink_ready(): await ctx.send("Still starting please wait a moment.") if guild_check: diff --git a/context.py b/context.py index 433fc9e..bee1485 100644 --- a/context.py +++ b/context.py @@ -8,16 +8,20 @@ from lavalink.models import DefaultPlayer if TYPE_CHECKING: from bot import TuneBot + from utils.classes import BaseCog class CustomContext(commands.Context): + bot: "TuneBot" + cog: "BaseCog" + @property def redis(self) -> Redis: return self.bot._redis_client @property def player(self) -> DefaultPlayer: - if hasattr(self.bot, "lavalink"): + if self.cog.is_lavalink_ready(): return self.bot.lavalink.player_manager.get(self.guild.id) raise CommandInvokeError("Lavalink is still starting up.") diff --git a/utils/classes.py b/utils/classes.py index 9cc95b2..7da2ad3 100644 --- a/utils/classes.py +++ b/utils/classes.py @@ -1,3 +1,4 @@ +import asyncio from typing import Dict from discord.ext.commands import Cog @@ -13,3 +14,9 @@ class BaseCog(Cog): for command in self.walk_commands(): if brief := slash_descriptions.get(command.qualified_name): command.brief = brief + + def is_lavalink_ready(self) -> bool: + return ( + hasattr(self.bot, "lavalink") + and len(self.bot.lavalink.node_manager.available_nodes) > 0 + ) diff --git a/utils/database.py b/utils/database.py index b8bcc06..05ac925 100644 --- a/utils/database.py +++ b/utils/database.py @@ -12,7 +12,7 @@ class AutoJoin: async def get_channels(redis: Redis) -> Dict[str, str]: # return all channels channels = await redis.hgetall(f"{redis_prefix}:autojoin") - return {key: value.split("-") for key, value in channels.items()} + return {key: value.split(":") for key, value in channels.items()} @staticmethod async def update_channel( @@ -21,7 +21,7 @@ class AutoJoin: await redis.hset( f"{redis_prefix}:autojoin", guild_id, - f"{voice_channel_id}-{text_channel_id}", + f"{voice_channel_id}:{text_channel_id}", ) @staticmethod From 9c55bcdc36daf54cb111ced897bf64fda77a5ea8 Mon Sep 17 00:00:00 2001 From: strNophix Date: Mon, 22 Nov 2021 11:11:17 +0100 Subject: [PATCH 4/5] Removed duplicate commands --- cogs/settings.py | 111 ----------------------------------------------- 1 file changed, 111 deletions(-) diff --git a/cogs/settings.py b/cogs/settings.py index 429ba86..d5f8d70 100644 --- a/cogs/settings.py +++ b/cogs/settings.py @@ -158,117 +158,6 @@ class SettingsCog(BaseCog, name="Settings"): await ctx.send(embed=embed) - @commands.is_owner() - @commands.group( - name="source", - aliases=["src"], - invoke_without_command=True, - slash_command=False, - hidden=True, - ) - async def source(self, ctx: CustomContext): - """Displays all possible options for the `source` command""" - prefix = self.bot.config["prefixes"][0] - embed = ctx.create_embed() - embed.title = "All options:" - embed.description = f"```{prefix}source list\n{prefix}source add \n{prefix}source remove \n{prefix}source sync```" - await ctx.send(embed=embed) - - @commands.is_owner() - @source.command(name="remove") - async def source_remove(self, ctx: CustomContext, source_url: str): - """Removes a source from the bot""" - if await PlaylistSource.remove(ctx.redis, source_url): - prefix = self.bot.config["prefixes"][0] - embed = ctx.create_embed() - embed.title = "Removed source succesfully" - embed.description = ( - f"Please use `{prefix}source sync` to persist these changes." - ) - await ctx.send(embed=embed) - return - - embed = ctx.create_embed() - embed.title = "Could not remove source, the specified source might not exist" - await ctx.send(embed=embed) - - @commands.is_owner() - @source.command(name="add") - async def source_add(self, ctx: CustomContext, source_url: str): - """ - Add's a source to the bot - - Supported sources: YouTube, SoundCloud, Bandcamp, Vimeo, Twitch and HTTP(S) URL's - """ - embed = ctx.create_embed() - embed.title = "Started processing source" - message = await ctx.send(embed=embed) - - query_result: Any = await self.bot.lavalink.get_tracks(source_url) - - if query_result["loadType"] == "LOAD_FAILED": - embed = ctx.create_embed() - embed.title = "The specified URL is not a valid source" - embed.description = f"Supported sources: YouTube, SoundCloud, Bandcamp, Vimeo, Twitch and HTTP(S) URL's" - if isinstance(message, Message): - await message.edit(embed=embed) - else: - await ctx.send(embed=embed) - return - - await PlaylistSource.add(ctx.redis, source_url) - track_urls = [str(track["info"]["uri"]) for track in query_result["tracks"]] - await Playlist.add_bulk(ctx.redis, track_urls) - - embed = ctx.create_embed() - embed.title = "Finished processing source" - embed.description = f"Added {len(track_urls)} tracks" - if isinstance(message, Message): - await message.edit(embed=embed) - else: - await ctx.send(embed=embed) - - @commands.is_owner() - @source.command(name="list", aliases=["ls"]) - async def source_list(self, ctx: CustomContext): - """Display a list of sources""" - # TODO: Implement pagination for sources - sources = await PlaylistSource.get_all(ctx.redis) - if len(sources) > 0: - description = "\n".join([f"[{source}]({source})" for source in sources]) - else: - description = "This bot has no sources yet" - - embed = ctx.create_embed() - embed.title = "All sources:" - embed.description = description - await ctx.send(embed=embed) - - @commands.is_owner() - @source.command(name="sync") - async def source_sync(self, ctx: CustomContext): - """Forcefully resyncs all sources""" - failed_sources: list[str] = [] - await Playlist.clear(ctx.redis) - sources = await PlaylistSource.get_all(ctx.redis) - for source_url in sources: - query_result: Any = await self.bot.lavalink.get_tracks(source_url) - if query_result["loadType"] == "LOAD_FAILED": - failed_sources.append(source_url) - continue - - track_urls = [str(track["info"]["uri"]) for track in query_result["tracks"]] - await Playlist.add_bulk(ctx.redis, track_urls) - - embed = ctx.create_embed() - embed.title = f"Finished sync ({len(failed_sources)} issues)" - if len(failed_sources) > 0: - embed.description = "\n".join( - [f"[{source_url}]({source_url})" for source_url in failed_sources] - ) - - await ctx.send(embed=embed) - def setup(bot: TuneBot): bot.add_cog(SettingsCog(bot)) From c45bd8ced314e237580ec6f602552c6091c22f2a Mon Sep 17 00:00:00 2001 From: strNophix Date: Mon, 22 Nov 2021 18:28:48 +0100 Subject: [PATCH 5/5] Added decorator for checking if eligable for managing sources --- cogs/settings.py | 11 ++++++----- config.json.sample | 1 + utils/decorators.py | 27 +++++++++++++++++++++++++++ 3 files changed, 34 insertions(+), 5 deletions(-) create mode 100644 utils/decorators.py diff --git a/cogs/settings.py b/cogs/settings.py index d5f8d70..64cb903 100644 --- a/cogs/settings.py +++ b/cogs/settings.py @@ -6,6 +6,7 @@ from discord.message import Message from bot import TuneBot from context import CustomContext from utils.classes import BaseCog +from utils.decorators import source_manager_only from utils.EmbedGenerator import EmbedGenerator @@ -47,7 +48,7 @@ class SettingsCog(BaseCog, name="Settings"): embed.title = f"AutoJoin disabled" await ctx.send(embed=embed) - @commands.is_owner() + @source_manager_only() @commands.group( name="source", aliases=["src"], @@ -63,7 +64,7 @@ class SettingsCog(BaseCog, name="Settings"): embed.description = f"```{prefix}source list\n{prefix}source add \n{prefix}source remove \n{prefix}source sync```" await ctx.send(embed=embed) - @commands.is_owner() + @source_manager_only() @source.command(name="remove") async def source_remove(self, ctx: CustomContext, source_url: str): """Removes a source from the bot""" @@ -81,7 +82,7 @@ class SettingsCog(BaseCog, name="Settings"): embed.title = "Could not remove source, the specified source might not exist" await ctx.send(embed=embed) - @commands.is_owner() + @source_manager_only() @source.command(name="add") async def source_add(self, ctx: CustomContext, source_url: str): """ @@ -117,7 +118,7 @@ class SettingsCog(BaseCog, name="Settings"): else: await ctx.send(embed=embed) - @commands.is_owner() + @source_manager_only() @source.command(name="list", aliases=["ls"]) async def source_list(self, ctx: CustomContext): """Display a list of sources""" @@ -133,7 +134,7 @@ class SettingsCog(BaseCog, name="Settings"): embed.description = description await ctx.send(embed=embed) - @commands.is_owner() + @source_manager_only() @source.command(name="sync") async def source_sync(self, ctx: CustomContext): """Forcefully resyncs all sources""" diff --git a/config.json.sample b/config.json.sample index 34e2c81..874b4a2 100644 --- a/config.json.sample +++ b/config.json.sample @@ -1,6 +1,7 @@ { "token": "", "owner_ids": [194545408960102400, 190875175460405249], + "manager_ids": [], "prefixes": ["ck!"], "redis_url": "", "redis_prefix": "", diff --git a/utils/decorators.py b/utils/decorators.py new file mode 100644 index 0000000..4851dfd --- /dev/null +++ b/utils/decorators.py @@ -0,0 +1,27 @@ +from typing import Callable +from typing import TYPE_CHECKING +from typing import TypeVar + +from discord.ext.commands import check +from discord.ext.commands.errors import NotOwner + +if TYPE_CHECKING: + from context import CustomContext + +T = TypeVar("T") + + +def source_manager_only() -> Callable[[T], T]: + """ + A :func:`.check` that checks if the person invoking this command is allowed to modify the radio sources. + """ + + async def predicate(ctx: "CustomContext") -> bool: + is_manager = ctx.author.id in ctx.bot.config["manager_ids"] + is_owner = await ctx.bot.is_owner(ctx.author) + if not is_manager and not is_owner: + raise NotOwner("You are not allowed to modify the sources.") + + return True + + return check(predicate)