diff --git a/bot.py b/bot.py index 3af0a2c..2bf0e1e 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 @@ -87,6 +111,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/information.py b/cogs/information.py index a8ba045..f20f2be 100644 --- a/cogs/information.py +++ b/cogs/information.py @@ -6,13 +6,11 @@ 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.classes import BaseCog -from utils.database import AutoJoin 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 282cb67..c480533 100644 --- a/cogs/music.py +++ b/cogs/music.py @@ -16,8 +16,6 @@ from lavalink.models import DefaultPlayer from bot import TuneBot from context import CustomContext from utils.classes import BaseCog -from utils.database import AutoJoin -from utils.database import Playlist from utils.EmbedGenerator import EmbedGenerator from utils.exceptions import EmbeddedCommandException @@ -76,25 +74,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 self.bot.global_autojoin.fetch_channels() 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) @@ -110,7 +94,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) @@ -153,7 +137,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/cogs/settings.py b/cogs/settings.py index 1dcaf0f..64cb903 100644 --- a/cogs/settings.py +++ b/cogs/settings.py @@ -6,9 +6,7 @@ from discord.message import Message from bot import TuneBot from context import CustomContext from utils.classes import BaseCog -from utils.database import AutoJoin -from utils.database import Playlist -from utils.database import PlaylistSource +from utils.decorators import source_manager_only from utils.EmbedGenerator import EmbedGenerator @@ -20,30 +18,37 @@ 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() + @source_manager_only() @commands.group( name="source", aliases=["src"], @@ -59,11 +64,11 @@ 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""" - if await PlaylistSource.remove(ctx.redis, source_url): + if await ctx.playlist_source.remove(source_url): prefix = self.bot.config["prefixes"][0] embed = ctx.create_embed() embed.title = "Removed source succesfully" @@ -77,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): """ @@ -101,9 +106,9 @@ class SettingsCog(BaseCog, name="Settings"): await ctx.send(embed=embed) return - await PlaylistSource.add(ctx.redis, source_url) + await ctx.playlist_source.add(source_url) track_urls = [str(track["info"]["uri"]) for track in query_result["tracks"]] - await Playlist.add_bulk(ctx.redis, track_urls) + await self.bot.global_playlist.add_tracks(track_urls) embed = ctx.create_embed() embed.title = "Finished processing source" @@ -113,12 +118,12 @@ 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""" # TODO: Implement pagination for sources - sources = await PlaylistSource.get_all(ctx.redis) + sources = await self.bot.global_playlist_source.fetch_sources() if len(sources) > 0: description = "\n".join([f"[{source}]({source})" for source in sources]) else: @@ -129,13 +134,13 @@ 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""" failed_sources: list[str] = [] - await Playlist.clear(ctx.redis) - sources = await PlaylistSource.get_all(ctx.redis) + 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": @@ -143,7 +148,7 @@ class SettingsCog(BaseCog, name="Settings"): continue track_urls = [str(track["info"]["uri"]) for track in query_result["tracks"]] - await Playlist.add_bulk(ctx.redis, track_urls) + await self.bot.global_playlist.add_tracks(track_urls) embed = ctx.create_embed() embed.title = f"Finished sync ({len(failed_sources)} issues)" 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/context.py b/context.py index 433fc9e..0020b0d 100644 --- a/context.py +++ b/context.py @@ -6,18 +6,27 @@ 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 + 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.") @@ -33,3 +42,15 @@ class CustomContext(commands.Context): 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/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..1e4a501 --- /dev/null +++ b/tunebot/redis/__init__.py @@ -0,0 +1,4 @@ +from tunebot.redis.entity import * # noreorder +from tunebot.redis.autojoin 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/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 deleted file mode 100644 index b8bcc06..0000000 --- a/utils/database.py +++ /dev/null @@ -1,57 +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) - - @staticmethod - async def add_bulk(redis: Redis, urls: list[str]): - await redis.sadd(f"{redis_prefix}:playlist", *urls) - - @staticmethod - async def clear(redis: Redis): - await redis.delete(f"{redis_prefix}:playlist") - - -class PlaylistSource: - @staticmethod - async def get_all(redis: Redis) -> list[str]: - return await redis.smembers(f"{redis_prefix}:sources") - - @staticmethod - async def add(redis: Redis, source_url: str): - await redis.sadd(f"{redis_prefix}:sources", source_url) - - @staticmethod - async def remove(redis: Redis, source_url: str) -> bool: - return await redis.srem(f"{redis_prefix}:sources", source_url) 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)