diff --git a/bot.py b/bot.py index 32cca39..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 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 436cce8..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 @@ -80,7 +78,7 @@ class Music(BaseCog): await asyncio.sleep(1) self.bot.lavalink.add_event_hook(self.track_hook) - redis_result = await AutoJoin.get_channels(self.bot._redis_client) + 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) @@ -96,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) diff --git a/cogs/settings.py b/cogs/settings.py index 1dcaf0f..429ba86 100644 --- a/cogs/settings.py +++ b/cogs/settings.py @@ -6,9 +6,6 @@ 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.EmbedGenerator import EmbedGenerator @@ -20,28 +17,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) @commands.is_owner() @commands.group( diff --git a/context.py b/context.py index bee1485..0020b0d 100644 --- a/context.py +++ b/context.py @@ -6,8 +6,13 @@ 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 @@ -37,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/database.py b/utils/database.py deleted file mode 100644 index 05ac925..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)