25 Commits
Author SHA1 Message Date
matthew bd0b04bf28 Merge pull request #44 from strNophix/lastfm-scrobbler
Plugins + lastfm scrobbler
2021-12-08 19:39:01 +01:00
niku 935383e7c5 Merge pull request #43 from strNophix/ensure-queue-size
fill_player_queue now recursively refetches until queue size is ensured
2021-12-07 06:41:16 +01:00
niku fdc9dc16ed Readjusted slash_command bot kwarg 2021-12-06 21:02:10 +01:00
niku 0dc2e0a7ca Added plugins key to the sample config 2021-12-06 20:22:40 +01:00
niku 2a24a4b5f0 Commit hoarding is a terrible practice 2021-12-06 20:14:24 +01:00
niku 0a14d9a93c fill_player_queue now recursively refetches until queue size is ensured 2021-12-03 21:24:57 +01:00
matthew 76bfd719c3 Merge pull request #40 from strNophix/source-manager-setting
Source manager decorator
2021-11-22 21:10:12 +01:00
niku c45bd8ced3 Added decorator for checking if eligable for managing sources 2021-11-22 18:28:48 +01:00
niku b37c3f533a Merge pull request #39 from strNophix/remove-duplicate-commands
Removed duplicate commands
2021-11-22 11:15:55 +01:00
niku 9c55bcdc36 Removed duplicate commands 2021-11-22 11:11:17 +01:00
niku 1495d214a1 Merge pull request #37 from strNophix/database-rewrite
database rewrite + autojoin
2021-11-22 10:53:23 +01:00
niku 4f65522bd8 Merge branch 'dev' into database-rewrite 2021-11-22 10:53:00 +01:00
niku 3e6ee9ecc9 Merge pull request #38 from strNophix/decoupling-lavalink
decoupled lavalink creation from cogs.music
2021-11-21 20:54:52 +01:00
niku 8504d6abfd decoupled lavalink creation from cogs.music 2021-11-20 00:37:17 +01:00
niku d2ecbbf6ae prevent circular import in tunebot.redis.__init__.py 2021-11-19 21:58:47 +01:00
niku adb3fe11a6 Merge remote-tracking branch 'upstream/dev' into database-rewrite 2021-11-19 21:39:10 +01:00
niku ebb07960bf database rewrite 2021-11-19 21:10:29 +01:00
matthew 1697002666 Merge pull request #35 from strNophix/source-management
Source management
2021-11-19 19:34:24 +01:00
niku 60f0bd9070 Made source commandgroup hidden 2021-11-18 23:03:19 +01:00
niku b8dc8b2b20 Added owner checks for source command 2021-11-18 23:01:04 +01:00
niku 998749f1ef Merge pull request #33 from strNophix/owner-id-fix
Critical bug with unsupplied owner_ids
2021-11-18 22:21:28 +01:00
niku 12eb3a70e1 Critical fix with owner_id 2021-11-18 22:17:12 +01:00
niku 316f3bacc9 Added sources functionality + pre-commit 2021-11-18 21:59:37 +01:00
matthew 2e007a6fb3 Merge pull request #28 from Matthww/fix_wlinfo
Hot fix for wlinfo
2021-11-06 22:28:39 +01:00
matthew 47d64c2cb3 Temp fix for wlinfo 2021-11-06 22:28:04 +01:00
30 changed files with 993 additions and 121 deletions
+1 -1
View File
@@ -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
+83 -17
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
+46
View File
@@ -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))
+84
View File
@@ -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)
+3 -3
View File
@@ -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)
+4 -3
View File
@@ -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)
+1
View File
@@ -0,0 +1 @@
from tunebot.abc import *
+115
View File
@@ -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",
)
+6
View File
@@ -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 *
+16
View File
@@ -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",)
+5
View File
@@ -0,0 +1,5 @@
class PluginInitFailed(Exception):
pass
__all__ = ("PluginInitFailed",)
+42
View File
@@ -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",)
+79
View File
@@ -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",)
+18
View File
@@ -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",)
+6
View File
@@ -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 *
+41
View File
@@ -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")
+49
View File
@@ -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")
+42
View File
@@ -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",)
+35
View File
@@ -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")
+18
View File
@@ -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)
+3 -2
View File
@@ -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
View File
@@ -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}")
-35
View File
@@ -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)
+27
View File
@@ -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
View File
@@ -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: