mirror of
https://github.com/Matthww/TuneBot.git
synced 2026-09-21 22:47:51 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bd0b04bf28 | ||
|
|
935383e7c5 | ||
|
|
fdc9dc16ed | ||
|
|
0dc2e0a7ca | ||
|
|
2a24a4b5f0 | ||
|
|
0a14d9a93c | ||
|
|
76bfd719c3 | ||
|
|
c45bd8ced3 | ||
|
|
b37c3f533a | ||
|
|
9c55bcdc36 | ||
|
|
1495d214a1 | ||
|
|
4f65522bd8 | ||
|
|
3e6ee9ecc9 | ||
|
|
8504d6abfd | ||
|
|
d2ecbbf6ae | ||
|
|
adb3fe11a6 | ||
|
|
ebb07960bf | ||
|
|
1697002666 | ||
|
|
60f0bd9070 | ||
|
|
b8dc8b2b20 | ||
|
|
998749f1ef | ||
|
|
12eb3a70e1 | ||
|
|
316f3bacc9 | ||
|
|
2e007a6fb3 | ||
|
|
47d64c2cb3 |
@@ -1,5 +1,5 @@
|
||||
default_language_version:
|
||||
python: python3.8
|
||||
python: python3.9
|
||||
repos:
|
||||
- repo: https://github.com/pre-commit/pre-commit-hooks
|
||||
rev: v4.0.1
|
||||
|
||||
@@ -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,28 +22,74 @@ from discord.ext.commands.errors import ExtensionNotFound
|
||||
from discord.ext.commands.errors import NoEntryPointError
|
||||
|
||||
from context import CustomContext
|
||||
from tunebot.plugins import FileSystemPluginLoader
|
||||
from tunebot.plugins import SimplePluginManager
|
||||
from tunebot.redis import GlobalRedisAutoJoin
|
||||
from tunebot.redis import GlobalRedisPlaylist
|
||||
from tunebot.redis import GlobalRedisPlaylistSource
|
||||
from tunebot.redis import GlobalRedisUtils
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tunebot import PluginManagerBase
|
||||
from tunebot import GlobalPlaylist
|
||||
from tunebot import GlobalPlaylistSource
|
||||
from tunebot import GlobalAutoJoin
|
||||
from tunebot import PluginLoaderBase
|
||||
from tunebot import PluginManagerBase
|
||||
from tunebot import GlobalUtils
|
||||
|
||||
ColorDict = dict[str, "Color"]
|
||||
|
||||
|
||||
class TuneBot(commands.Bot):
|
||||
lavalink: lavalink.Client
|
||||
invite_link: str
|
||||
invite_link: str = ""
|
||||
initial_cog_names: list[str]
|
||||
colors: ColorDict
|
||||
|
||||
global_autojoin: "GlobalAutoJoin"
|
||||
global_playlist: "GlobalPlaylist"
|
||||
global_playlist_source: "GlobalPlaylistSource"
|
||||
global_utils: "GlobalUtils"
|
||||
|
||||
plugin_loader: "PluginLoaderBase"
|
||||
plugin_manager: "PluginManagerBase"
|
||||
|
||||
def __init__(self, config: Dict[Any, Any]):
|
||||
intents = discord.Intents(
|
||||
voice_states=True, guild_messages=True, guilds=True, messages=True
|
||||
voice_states=True,
|
||||
guild_messages=True,
|
||||
guilds=True,
|
||||
messages=True,
|
||||
members=True,
|
||||
)
|
||||
|
||||
self.rpc_is_help_message = True
|
||||
self.update_status.start()
|
||||
|
||||
self.config = config
|
||||
self.initial_cog_names: List[str] = self.config.get("cogs", [])
|
||||
self.colors: Dict[str, Color] = self.process_colours(config.get("colors", []))
|
||||
|
||||
self.initial_cog_names = self.config.get("cogs", [])
|
||||
self.colors = 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.invite_link: str = ""
|
||||
|
||||
self.global_autojoin = GlobalRedisAutoJoin(
|
||||
self._redis_client, self.redis_prefix
|
||||
)
|
||||
self.global_playlist = GlobalRedisPlaylist(
|
||||
self._redis_client, self.redis_prefix
|
||||
)
|
||||
self.global_playlist_source = GlobalRedisPlaylistSource(
|
||||
self._redis_client, self.redis_prefix
|
||||
)
|
||||
self.global_utils = GlobalRedisUtils(self._redis_client, self.redis_prefix)
|
||||
|
||||
self.plugin_loader = FileSystemPluginLoader(self)
|
||||
self.plugin_manager = SimplePluginManager(self)
|
||||
|
||||
slash_guilds = None
|
||||
if len(self.config["slash_command_guilds"]) > 0:
|
||||
@@ -50,6 +97,7 @@ class TuneBot(commands.Bot):
|
||||
|
||||
super().__init__(
|
||||
command_prefix=self.prefix_callable,
|
||||
owner_ids=self.config["owner_ids"],
|
||||
description=self.config["info"]["description"],
|
||||
case_insensitive=False,
|
||||
fetch_offline_members=False,
|
||||
@@ -61,23 +109,24 @@ class TuneBot(commands.Bot):
|
||||
self.loop.create_task(self.async_init())
|
||||
|
||||
async def async_init(self):
|
||||
await self.load_cogs(self.initial_cog_names)
|
||||
self.init_plugins()
|
||||
self.load_cogs(self.initial_cog_names)
|
||||
|
||||
async def prefix_callable(self, _, msg: Message) -> List[str]:
|
||||
return commands.when_mentioned_or(*self.config["prefixes"])(self, msg)
|
||||
|
||||
async def load_cogs(self, cog_names: Sequence[str]):
|
||||
def load_cogs(self, cog_names: Sequence[str]):
|
||||
for cog in cog_names:
|
||||
try:
|
||||
self.load_extension(cog)
|
||||
print(f"Succesfully loaded extension {cog}.")
|
||||
print(f"[✓] loaded extension: {cog}.")
|
||||
except (
|
||||
ExtensionNotFound,
|
||||
ExtensionAlreadyLoaded,
|
||||
NoEntryPointError,
|
||||
ExtensionFailed,
|
||||
) as e:
|
||||
print(f"Failed to load extension {cog}.\n\t{e}", file=sys.stderr)
|
||||
print(f"[x] failed loading extension: {cog}.\n\t{e}", file=sys.stderr)
|
||||
|
||||
async def on_ready(self):
|
||||
self.invite_link = f"https://discord.com/oauth2/authorize?client_id={self.user.id}&permissions=3230720&scope=bot%20applications.commands"
|
||||
@@ -86,7 +135,13 @@ class TuneBot(commands.Bot):
|
||||
print(f"Version: {discord.__version__}")
|
||||
print(f"Invite: {self.invite_link}")
|
||||
|
||||
def process_colours(self, colors: Dict[str, str]) -> Dict[str, Color]:
|
||||
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]) -> ColorDict:
|
||||
colour_dict: Dict[str, Color] = {}
|
||||
for name, color in colors.items():
|
||||
colour_dict[name] = Color(int(color, 16))
|
||||
@@ -95,6 +150,19 @@ class TuneBot(commands.Bot):
|
||||
async def get_context(self, message: Message, *, cls=CustomContext):
|
||||
return await super().get_context(message, cls=cls)
|
||||
|
||||
def init_plugins(self):
|
||||
for plug_id, plug_conf in self.config["plugins"].items():
|
||||
if not plug_conf.get("enabled"):
|
||||
continue
|
||||
|
||||
try:
|
||||
plugin = self.plugin_loader.load_plugin(plug_conf)
|
||||
self.plugin_manager.enable_plugin(plug_id, plugin)
|
||||
print(f"[✓] loaded plugin: {plug_id}")
|
||||
except Exception as e:
|
||||
self.plugin_manager.remove_plugin(plug_id)
|
||||
print(f"[x] failed loading plugin: {plug_id}\n{e}")
|
||||
|
||||
@tasks.loop(seconds=30)
|
||||
async def update_status(self):
|
||||
await self.wait_until_ready()
|
||||
@@ -110,13 +178,6 @@ class TuneBot(commands.Bot):
|
||||
await self.change_presence(activity=activity)
|
||||
|
||||
|
||||
config_path = "config.json"
|
||||
if len(sys.argv) > 1:
|
||||
config_path = sys.argv[1]
|
||||
|
||||
config = json.load(open(config_path, "r", encoding="utf-8"))
|
||||
redis_prefix = config["redis_prefix"]
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
import uvloop
|
||||
@@ -126,5 +187,10 @@ if __name__ == "__main__":
|
||||
except ModuleNotFoundError:
|
||||
pass
|
||||
|
||||
config_path = "config.json"
|
||||
if len(sys.argv) > 1:
|
||||
config_path = sys.argv[1]
|
||||
|
||||
config = json.load(open(config_path, "r", encoding="utf-8"))
|
||||
token = config.pop("token")
|
||||
TuneBot(config).run(token, reconnect=True)
|
||||
|
||||
+3
-6
@@ -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().identifier.__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)
|
||||
|
||||
+47
-32
@@ -1,27 +1,28 @@
|
||||
import asyncio
|
||||
import datetime
|
||||
import re
|
||||
from typing import Any
|
||||
from typing import Optional
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
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 tunebot.plugins import ServiceEvent
|
||||
from utils.classes import BaseCog
|
||||
from utils.EmbedGenerator import EmbedGenerator
|
||||
from utils.exceptions import EmbeddedCommandException
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from discord import VoiceChannel
|
||||
|
||||
url_rx = re.compile(r"https?://(?:www\.)?.+")
|
||||
|
||||
@@ -78,25 +79,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)
|
||||
@@ -112,11 +99,12 @@ 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)
|
||||
# Get the results for the query from Lavalink.
|
||||
queries = await self.bot.global_playlist.pick_random(buffer)
|
||||
failed_queries: list[str] = []
|
||||
for query in queries:
|
||||
result = await player.node.get_tracks(query)
|
||||
if not result or not result["tracks"]:
|
||||
failed_queries.append(query)
|
||||
continue
|
||||
|
||||
track = lavalink.models.AudioTrack(
|
||||
@@ -124,6 +112,10 @@ class Music(BaseCog):
|
||||
)
|
||||
player.add(requester=self.bot.user.id, track=track)
|
||||
|
||||
if len(failed_queries) > 0:
|
||||
await self.bot.global_playlist.remove_tracks(failed_queries)
|
||||
await self.fill_player_queue(player, len(failed_queries))
|
||||
|
||||
async def create_track_embed(self, track: AudioTrack) -> Embed:
|
||||
embed_color = self.bot.colors["embed"]
|
||||
embed = discord.Embed(
|
||||
@@ -155,7 +147,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:
|
||||
@@ -230,13 +222,36 @@ class Music(BaseCog):
|
||||
guild = self.bot.get_guild(guild_id)
|
||||
await guild.voice_client.disconnect(force=True)
|
||||
elif isinstance(event, lavalink.events.TrackStartEvent):
|
||||
channel_id = int(event.player.fetch("channel"))
|
||||
channel: TextChannel = self.bot.get_channel(channel_id)
|
||||
embed = await self.create_track_embed(event.player.current)
|
||||
await channel.send(embed=embed)
|
||||
if channel := event.player.fetch("channel"):
|
||||
channel_id = int(channel)
|
||||
channel: TextChannel = self.bot.get_channel(channel_id)
|
||||
embed = await self.create_track_embed(event.player.current)
|
||||
await channel.send(embed=embed)
|
||||
elif isinstance(event, lavalink.events.TrackEndEvent):
|
||||
await self.fill_player_queue(event.player, 1)
|
||||
|
||||
if event.reason == "FINISHED":
|
||||
if event.player.channel_id:
|
||||
channel_id = int(event.player.channel_id)
|
||||
voice_channel: "VoiceChannel" = self.bot.get_channel(channel_id)
|
||||
in_voice: list[int] = []
|
||||
|
||||
for member in voice_channel.members:
|
||||
if member.bot:
|
||||
continue
|
||||
|
||||
in_voice.append(member.id)
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"last_track": event.track,
|
||||
"in_voice": in_voice,
|
||||
}
|
||||
s = ServiceEvent.TRACK_ENDED
|
||||
await self.bot.plugin_manager.dispatch(s, payload)
|
||||
else:
|
||||
fmt = f"Failed dispatching for TrackEnd event, missing channel_id on player"
|
||||
print(fmt)
|
||||
|
||||
@commands.command(name="connect", aliases=["p", "play", "join"])
|
||||
async def play(self, ctx: CustomContext):
|
||||
"""Start the radio"""
|
||||
|
||||
+134
-17
@@ -1,13 +1,12 @@
|
||||
from typing import Any
|
||||
|
||||
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 discord.message import Message
|
||||
|
||||
from bot import TuneBot
|
||||
from context import CustomContext
|
||||
from utils.database import AutoJoin
|
||||
from utils.classes import BaseCog
|
||||
from utils.decorators import source_manager_only
|
||||
from utils.EmbedGenerator import EmbedGenerator
|
||||
|
||||
|
||||
@@ -19,28 +18,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)
|
||||
|
||||
@source_manager_only()
|
||||
@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 <url>\n{prefix}source remove <url>\n{prefix}source sync```"
|
||||
await ctx.send(embed=embed)
|
||||
|
||||
@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 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)
|
||||
|
||||
@source_manager_only()
|
||||
@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)
|
||||
|
||||
@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 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)
|
||||
|
||||
@source_manager_only()
|
||||
@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):
|
||||
|
||||
+15
-1
@@ -1,6 +1,7 @@
|
||||
{
|
||||
"token": "",
|
||||
"owner_ids": [194545408960102400, 190875175460405249],
|
||||
"manager_ids": [],
|
||||
"prefixes": ["ck!"],
|
||||
"redis_url": "",
|
||||
"redis_prefix": "",
|
||||
@@ -21,5 +22,18 @@
|
||||
"cogs": ["cogs.owner", "cogs.settings", "cogs.information", "cogs.music"],
|
||||
"slash_command_guilds": [],
|
||||
"queue_buffer_size": 5,
|
||||
"slash_descriptions": {}
|
||||
"slash_descriptions": {},
|
||||
"plugins": {
|
||||
"lastfm-scrobbler": {
|
||||
"services": ["plugins.lastfm_scrobbler.service"],
|
||||
"cogs": ["plugins.lastfm_scrobbler.cog"],
|
||||
"config": {
|
||||
"lastfm_api_key": "",
|
||||
"lastfm_api_secret": "",
|
||||
"session_key_table": "",
|
||||
"auth_url": ""
|
||||
},
|
||||
"enabled": false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+47
-1
@@ -1,17 +1,63 @@
|
||||
from typing import Any
|
||||
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
|
||||
from utils.classes import BaseCog
|
||||
from tunebot.plugins import ServiceEvent
|
||||
|
||||
AnyDict = dict[Any, Any]
|
||||
|
||||
|
||||
class CustomContext(commands.Context):
|
||||
bot: "TuneBot"
|
||||
cog: "BaseCog"
|
||||
|
||||
async def dispatch(self, event: "ServiceEvent", payload: AnyDict):
|
||||
await self.bot.plugin_manager.dispatch(event, payload)
|
||||
|
||||
@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.")
|
||||
|
||||
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
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from discord.ext import commands
|
||||
|
||||
from context import CustomContext
|
||||
from utils.classes import PluginCog
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from bot import TuneBot
|
||||
|
||||
|
||||
class LastFMScrobblerCog(PluginCog, name="LastFMScrobbler"):
|
||||
def __init__(self, bot: "TuneBot") -> None:
|
||||
super().__init__(bot)
|
||||
self.plugin = self.get_plugin_instance("lastfm-scrobbler")
|
||||
|
||||
@commands.group(name="lfm", invoke_without_command=True)
|
||||
@commands.cooldown(rate=1, per=5, type=commands.BucketType.user)
|
||||
async def lfm(self, ctx: CustomContext):
|
||||
"""Enable/Disable LastFM Scrobbling"""
|
||||
embed = ctx.create_embed()
|
||||
embed.title = f"LastFM scrobbler usage:"
|
||||
embed.description = f"```{ctx.prefix}lfm enable\n{ctx.prefix}lfm del```"
|
||||
await ctx.send(embed=embed)
|
||||
|
||||
@lfm.command(name="enable", aliases=["set"])
|
||||
async def lfm_aut(self, ctx: CustomContext):
|
||||
"""Enable LastFM Scrobbling"""
|
||||
embed = ctx.create_embed()
|
||||
embed.title = f"Start scrobbling"
|
||||
auth_url = self.plugin.config["auth_url"] + "?user_id=" + str(ctx.author.id)
|
||||
embed.description = f"[Login]({auth_url})"
|
||||
await ctx.send(embed=embed)
|
||||
|
||||
@lfm.command(name="delete", aliases=["del"])
|
||||
async def lfm_del(self, ctx: CustomContext):
|
||||
"""Disable LastFM Scrobbling"""
|
||||
db_key = self.plugin.config["session_key_table"]
|
||||
await self.bot.global_utils.raw_table_del_entry(db_key, [ctx.author.id])
|
||||
embed = ctx.create_embed()
|
||||
embed.title = f"Stopped scrobbling"
|
||||
await ctx.send(embed=embed)
|
||||
|
||||
|
||||
def setup(bot: "TuneBot"):
|
||||
bot.add_cog(LastFMScrobblerCog(bot))
|
||||
@@ -0,0 +1,84 @@
|
||||
import hashlib
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import aiohttp
|
||||
from lavalink.models import AudioTrack
|
||||
|
||||
from tunebot.plugins import ServiceEvent
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from bot import TuneBot
|
||||
|
||||
AnyDict = dict[Any, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class Track:
|
||||
name: str
|
||||
artist: str
|
||||
|
||||
|
||||
class LastFMScrobbler:
|
||||
api_url = "http://ws.audioscrobbler.com/2.0/"
|
||||
|
||||
def __init__(self, bot: "TuneBot", config: AnyDict) -> None:
|
||||
self.bot = bot
|
||||
self.plug_conf = config
|
||||
|
||||
async def on_dispatch(self, event: "ServiceEvent", payload: AnyDict):
|
||||
if event == ServiceEvent.TRACK_ENDED:
|
||||
_track: AudioTrack = payload.get("last_track", None)
|
||||
if not _track:
|
||||
print("LastFMScrobbler: Could not find required field `last_track`")
|
||||
return
|
||||
|
||||
# TODO: expand parsing options
|
||||
artist, track_name = _track.title.split(" - ", 1)
|
||||
track = Track(track_name, artist)
|
||||
|
||||
in_voice = payload.get("in_voice", [])
|
||||
for session_key in await self.fetch_lastfm_sessions(in_voice):
|
||||
if not session_key:
|
||||
continue
|
||||
|
||||
await self.scrobble(session_key, track)
|
||||
|
||||
async def fetch_lastfm_sessions(self, user_ids: list[int]) -> list[str]:
|
||||
table_name = self.plug_conf["config"]["session_key_table"]
|
||||
result: list[str] = await self.bot.global_utils.raw_table_lookup(
|
||||
table_name, user_ids
|
||||
)
|
||||
return result
|
||||
|
||||
async def scrobble(self, session_key: str, track: Track):
|
||||
params: AnyDict = {
|
||||
"method": "track.scrobble",
|
||||
"timestamp": str(int(time.time() - 30)),
|
||||
"track": track.name,
|
||||
"artist": track.artist,
|
||||
"sk": session_key,
|
||||
}
|
||||
resp = await self.lastfm_request(params)
|
||||
if resp.status != 200:
|
||||
fmt = f"Failed to scrobble for user {session_key} on track: {track.artist} - {track.name}"
|
||||
print(fmt)
|
||||
|
||||
async def lastfm_request(self, params: AnyDict) -> aiohttp.ClientResponse:
|
||||
params["api_key"] = self.plug_conf["config"]["lastfm_api_key"]
|
||||
params = {key: params[key] for key in sorted(params)}
|
||||
|
||||
secret = self.plug_conf["config"]["lastfm_api_secret"]
|
||||
sig_str = "".join(key + params[key] for key in params.keys()) + secret
|
||||
params["api_sig"] = hashlib.md5(sig_str.encode("utf8")).hexdigest()
|
||||
|
||||
params["format"] = "json"
|
||||
async with aiohttp.ClientSession() as sess:
|
||||
async with sess.post(self.api_url, params=params) as resp:
|
||||
return resp
|
||||
|
||||
|
||||
def setup(bot: "TuneBot", config: AnyDict):
|
||||
return LastFMScrobbler(bot, config)
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
print(yt_url)
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
from tunebot.abc import *
|
||||
+115
@@ -0,0 +1,115 @@
|
||||
from abc import ABC
|
||||
from abc import abstractmethod
|
||||
from typing import Any
|
||||
from typing import Optional
|
||||
from typing import Protocol
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import Union
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tunebot.plugins import ServiceEvent
|
||||
from discord.ext.commands import Cog
|
||||
|
||||
AnyDict = dict[Any, Any]
|
||||
|
||||
|
||||
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
|
||||
|
||||
@abstractmethod
|
||||
async def remove_tracks(self, track_urls: list[str]):
|
||||
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
|
||||
|
||||
|
||||
class ServiceBase(Protocol):
|
||||
async def on_dispatch(self, event: "ServiceEvent", payload: AnyDict):
|
||||
...
|
||||
|
||||
|
||||
class PluginManagerBase(Protocol):
|
||||
def get_plugin(self, name: str) -> Union["BasePluginInstance", None]:
|
||||
...
|
||||
|
||||
def remove_plugin(self, name: str):
|
||||
...
|
||||
|
||||
def enable_plugin(self, name: str, plugin: "BasePluginInstance"):
|
||||
...
|
||||
|
||||
async def dispatch(self, event: "ServiceEvent", payload: AnyDict = {}):
|
||||
...
|
||||
|
||||
|
||||
class PluginLoaderBase(Protocol):
|
||||
def load_plugin(self, plug_conf: AnyDict) -> "BasePluginInstance":
|
||||
...
|
||||
|
||||
|
||||
class GlobalUtils(Protocol):
|
||||
async def raw_table_lookup(self, table_name: str, keys: list[Any]) -> list[Any]:
|
||||
...
|
||||
|
||||
async def raw_table_del_entry(self, table_name: str, keys: list[Any]):
|
||||
...
|
||||
|
||||
|
||||
class BasePluginInstance(Protocol):
|
||||
config: AnyDict
|
||||
services: list["ServiceBase"]
|
||||
cogs: list[str]
|
||||
|
||||
|
||||
__all__ = (
|
||||
"GlobalPlaylistSource",
|
||||
"PlaylistSource",
|
||||
"GlobalPlaylist",
|
||||
"GlobalAutoJoin",
|
||||
"AutoJoin",
|
||||
"ServiceBase",
|
||||
"PluginManagerBase",
|
||||
"PluginLoaderBase",
|
||||
"GlobalUtils",
|
||||
"BasePluginInstance",
|
||||
)
|
||||
@@ -0,0 +1,6 @@
|
||||
# noreorder
|
||||
from tunebot.plugins.exceptions import *
|
||||
from tunebot.plugins.events import *
|
||||
from tunebot.plugins.plugin import *
|
||||
from tunebot.plugins.loader import *
|
||||
from tunebot.plugins.manager import *
|
||||
@@ -0,0 +1,16 @@
|
||||
from enum import auto
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class ServiceEvent(Enum):
|
||||
"""
|
||||
This Enum contains all possible events that can be dispatched to services
|
||||
|
||||
Args:
|
||||
Enum ([type]): [description]
|
||||
"""
|
||||
|
||||
TRACK_ENDED = auto()
|
||||
|
||||
|
||||
__all__ = ("ServiceEvent",)
|
||||
@@ -0,0 +1,5 @@
|
||||
class PluginInitFailed(Exception):
|
||||
pass
|
||||
|
||||
|
||||
__all__ = ("PluginInitFailed",)
|
||||
@@ -0,0 +1,42 @@
|
||||
import importlib
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from tunebot.plugins import PluginInstance
|
||||
from tunebot.plugins.exceptions import PluginInitFailed
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tunebot import BasePluginInstance
|
||||
from tunebot import ServiceBase
|
||||
from bot import TuneBot
|
||||
|
||||
AnyDict = dict[Any, Any]
|
||||
|
||||
|
||||
class FileSystemPluginLoader:
|
||||
def __init__(self, bot: "TuneBot") -> None:
|
||||
self.bot = bot
|
||||
|
||||
def load_plugin(self, plug_conf: AnyDict) -> "BasePluginInstance":
|
||||
"""
|
||||
Loads a plugin from configuration and returns a sequence of Services
|
||||
|
||||
Raises:
|
||||
PluginInitFailed: [description]
|
||||
|
||||
Returns:
|
||||
tuple[list["ServiceBase"], list[str]]: [description]
|
||||
"""
|
||||
services: list["ServiceBase"] = []
|
||||
for service_location in plug_conf["services"]:
|
||||
module = importlib.import_module(service_location)
|
||||
|
||||
if not hasattr(module, "setup"):
|
||||
raise PluginInitFailed('Failed to find setup() for "{location}"')
|
||||
|
||||
services.append(module.setup(self.bot, plug_conf))
|
||||
|
||||
return PluginInstance(plug_conf["config"], services, plug_conf["cogs"])
|
||||
|
||||
|
||||
__all__ = ("FileSystemPluginLoader",)
|
||||
@@ -0,0 +1,79 @@
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import Union
|
||||
|
||||
from discord.ext.commands.errors import ExtensionAlreadyLoaded
|
||||
from discord.ext.commands.errors import ExtensionFailed
|
||||
from discord.ext.commands.errors import ExtensionNotFound
|
||||
from discord.ext.commands.errors import NoEntryPointError
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tunebot.abc import BasePluginInstance
|
||||
from tunebot.plugins import ServiceEvent
|
||||
from bot import TuneBot
|
||||
|
||||
AnyDict = dict[Any, Any]
|
||||
|
||||
|
||||
class SimplePluginManager:
|
||||
_plugins: dict[str, "BasePluginInstance"] = {}
|
||||
|
||||
def __init__(self, bot: "TuneBot") -> None:
|
||||
self.bot = bot
|
||||
|
||||
def get_plugin(self, name: str) -> Union["BasePluginInstance", None]:
|
||||
"""
|
||||
Retrieves the corresponding `BasePluginInstance` if it exists
|
||||
|
||||
Returns:
|
||||
Union["BasePluginInstance", None]: [description]
|
||||
"""
|
||||
return self._plugins.get(name)
|
||||
|
||||
def remove_plugin(self, name: str):
|
||||
"""
|
||||
Unloads/Removes all components related to the `BasePluginInstance`
|
||||
|
||||
Args:
|
||||
name (str): [description]
|
||||
"""
|
||||
plugin = self.get_plugin(name)
|
||||
|
||||
if not plugin:
|
||||
return
|
||||
|
||||
for cog in plugin.cogs:
|
||||
try:
|
||||
self.bot.unload_extension(cog)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
del self._plugins[name]
|
||||
|
||||
def enable_plugin(self, name: str, plugin: "BasePluginInstance"):
|
||||
"""
|
||||
Loads/Activates all cogs/services included within the plugin
|
||||
|
||||
Args:
|
||||
plugin_name (str): [description]
|
||||
services (list[): [description]
|
||||
cog_names (list[str]): [description]
|
||||
"""
|
||||
self._plugins[name] = plugin
|
||||
for cog_name in plugin.cogs:
|
||||
self.bot.load_extension(cog_name)
|
||||
|
||||
async def dispatch(self, event: "ServiceEvent", payload: AnyDict = {}):
|
||||
"""
|
||||
Dispatches an event to all registered services
|
||||
|
||||
Args:
|
||||
event (ServiceEvent): [description]
|
||||
payload (AnyDict, optional): [description]. Defaults to {}.
|
||||
"""
|
||||
for plugin in self._plugins.values():
|
||||
for service in plugin.services:
|
||||
await service.on_dispatch(event, payload)
|
||||
|
||||
|
||||
__all__ = ("SimplePluginManager",)
|
||||
@@ -0,0 +1,18 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tunebot import ServiceBase
|
||||
|
||||
AnyDict = dict[Any, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class PluginInstance:
|
||||
config: AnyDict
|
||||
services: list["ServiceBase"]
|
||||
cogs: list[str]
|
||||
|
||||
|
||||
__all__ = ("PluginInstance",)
|
||||
@@ -0,0 +1,6 @@
|
||||
# noreorder
|
||||
from tunebot.redis.entity import *
|
||||
from tunebot.redis.utils import *
|
||||
from tunebot.redis.autojoin import *
|
||||
from tunebot.redis.playlist import *
|
||||
from tunebot.redis.playlist_source import *
|
||||
@@ -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")
|
||||
@@ -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")
|
||||
@@ -0,0 +1,42 @@
|
||||
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"))
|
||||
|
||||
async def remove_tracks(self, track_urls: list[str]):
|
||||
"""
|
||||
Removes a single track from the playlist
|
||||
|
||||
Args:
|
||||
track_url (str): [description]
|
||||
"""
|
||||
await self.redis.srem(self.key("playlist"), *track_urls)
|
||||
|
||||
|
||||
__all__ = ("GlobalRedisPlaylist",)
|
||||
@@ -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")
|
||||
@@ -0,0 +1,18 @@
|
||||
from typing import Any
|
||||
|
||||
from tunebot.redis import RedisBotEntity
|
||||
|
||||
|
||||
class GlobalRedisUtils(RedisBotEntity):
|
||||
async def raw_table_lookup(self, table_name: str, keys: list[Any]) -> list[Any]:
|
||||
if len(keys) == 0:
|
||||
return []
|
||||
|
||||
result: list[Any] = await self.redis.hmget(table_name, keys)
|
||||
return result
|
||||
|
||||
async def raw_table_del_entry(self, table_name: str, keys: list[Any]):
|
||||
if len(keys) == 0:
|
||||
return
|
||||
|
||||
await self.redis.hdel(table_name, *keys)
|
||||
@@ -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:
|
||||
|
||||
+21
-2
@@ -1,13 +1,32 @@
|
||||
from typing import Dict
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from discord.ext.commands import Cog
|
||||
from bot import TuneBot
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tunebot import BasePluginInstance
|
||||
from bot import TuneBot
|
||||
|
||||
|
||||
class BaseCog(Cog):
|
||||
def __init__(self, bot: TuneBot) -> None:
|
||||
def __init__(self, bot: "TuneBot") -> None:
|
||||
self.bot = bot
|
||||
|
||||
slash_descriptions: Dict[str, str] = self.bot.config["slash_descriptions"]
|
||||
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
|
||||
)
|
||||
|
||||
|
||||
class PluginCog(BaseCog):
|
||||
def get_plugin_instance(self, name: str) -> "BasePluginInstance":
|
||||
if plugin := self.bot.plugin_manager.get_plugin(name):
|
||||
return plugin
|
||||
|
||||
raise KeyError(f"Failed to retrieve plugin instance with name: {name}")
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
+2
-1
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user