diff --git a/koala/cogs/react_for_role/__init__.py b/koala/cogs/react_for_role/__init__.py index f38c35dd..e8443764 100644 --- a/koala/cogs/react_for_role/__init__.py +++ b/koala/cogs/react_for_role/__init__.py @@ -1,2 +1,7 @@ -from . import utils, db, models -from .cog import ReactForRole, setup +from . import utils, db, models, cog, core, api +from .cog import ReactForRole + + +async def setup(bot): + await cog.setup(bot) + api.setup(bot) diff --git a/koala/cogs/react_for_role/api.py b/koala/cogs/react_for_role/api.py new file mode 100644 index 00000000..d915d71b --- /dev/null +++ b/koala/cogs/react_for_role/api.py @@ -0,0 +1,240 @@ +# Futures +# Built-in/Generic Imports +# Libs +from typing import List + +import discord +from aiohttp import web +from discord.ext.commands import Bot + +import koalabot +from koala.rest.api import parse_request + +# Own modules +from . import core +from .dto import ReactRole +from .log import logger +from ... import colours + +# Constants +RFR_ENDPOINT = 'react-for-role' + +MESSAGE = 'message' +REQUIRED_ROLES = 'required-roles' + + +class RfrEndpoint: + _bot: koalabot.KoalaBot + """ + The API endpoints for BaseCog + """ + + def __init__(self, bot): + self._bot = bot + + def register(self, app): + """ + Register the routes for the given application + :param app: The aiohttp.web.Application (likely of the sub app) + :return: app + """ + app.add_routes([web.post('/{}'.format(MESSAGE), self.post_message), + web.get('/{}'.format(MESSAGE), self.get_message), + web.put('/{}'.format(MESSAGE), self.put_message), + web.patch('/{}'.format(MESSAGE), self.patch_message), + web.delete('/{}'.format(MESSAGE), self.delete_message), + web.put('/{}'.format(REQUIRED_ROLES), self.put_required_roles), + web.get('/{}'.format(REQUIRED_ROLES), self.get_required_roles)]) + return app + + @parse_request + async def post_message(self, + guild_id: int, + channel_id: int, + title: str, + description: str = "", + colour: str = colours.KOALA_GREEN.__str__(), + thumbnail: str = None, + inline: bool = None, + roles: List[dict] = None + ): + """ + Create a React For Role message + + :param guild_id: Guild ID of RFR message + :param channel_id: Channel ID of RFR message + :param title: Title of RFR message + :param description: Description of RFR message + :param colour: Hex colour code of RFR message + :param thumbnail: thumbnail URL + :param inline: fields should be inline + :param roles: roles for RFR message + :return: + """ + guild = self._bot.get_guild(guild_id) + if roles is not None: + roles = [ReactRole(r["emoji"], r["role_id"]).to_tuple(guild) for r in roles] + + return await core.create_rfr_message(bot=self._bot, + guild_id=guild_id, + channel_id=channel_id, + title=title, + description=description, + colour=discord.Colour.from_str(colour), + thumbnail=thumbnail, + inline=inline, + roles=roles) + + @parse_request + async def get_message(self, + message_id: int, + guild_id: int, + channel_id: int + ): + """ + Get a React For Role message + + :param message_id: Message ID of RFR message + :param guild_id: Guild ID of RFR message + :param channel_id: Channel ID of RFR message + :return: + """ + return await core.get_rfr_message_dto(self._bot, int(message_id), int(guild_id), int(channel_id)) + + @parse_request + async def put_message(self, + message_id: int, + guild_id: int, + channel_id: int, + title: str, + description: str, + colour: str, + thumbnail: str, + inline: bool, + roles: List[dict] + ): + """ + Edit a React For Role message + + :param message_id: Message ID of RFR message + :param guild_id: Guild ID of RFR message + :param channel_id: Channel ID of RFR message + :param title: Title of RFR message + :param description: Description of RFR message + :param colour: Hex colour code of RFR message + :param thumbnail: thumbnail URL + :param inline: fields should be inline + :param roles: roles for RFR message + :return: + """ + guild = self._bot.get_guild(guild_id) + if roles is not None: + roles = [ReactRole(r["emoji"], r["role_id"]).to_tuple(guild) for r in roles] + return await core.update_rfr_message(bot=self._bot, + message_id=message_id, + guild_id=guild_id, + channel_id=channel_id, + title=title, + description=description, + colour=discord.Colour.from_str(colour), + thumbnail=thumbnail, + inline=inline, + roles=roles) + + @parse_request + async def patch_message(self, + message_id: int, + guild_id: int, + channel_id: int, + title: str = None, + description: str = None, + colour: str = None, + thumbnail: str = None, + inline: bool = None, + roles: List[dict] = None + ): + """ + Edit a React For Role message + + :param message_id: Message ID of RFR message + :param guild_id: Guild ID of RFR message + :param channel_id: Channel ID of RFR message + :param title: Title of RFR message + :param description: Description of RFR message + :param colour: Hex colour code of RFR message + :param thumbnail: thumbnail URL + :param inline: fields should be inline + :param roles: roles for RFR message + :return: + """ + guild = self._bot.get_guild(guild_id) + if roles is not None: + roles = [ReactRole(r["emoji"], r["role_id"]).to_tuple(guild) for r in roles] + + if colour is not None: + colour = discord.Colour.from_str(colour) + + return await core.update_rfr_message(bot=self._bot, + message_id=message_id, + guild_id=guild_id, + channel_id=channel_id, + title=title, + description=description, + colour=colour, + thumbnail=thumbnail, + inline=inline, + roles=roles) + + @parse_request + async def delete_message(self, + message_id: int, + guild_id: int, + channel_id: int + ): + """ + Delete a React For Role message + + :param message_id: Message ID of RFR message + :param guild_id: Guild ID of RFR message + :param channel_id: Channel ID of RFR message + :return: + """ + await core.delete_rfr_message(self._bot, int(message_id), int(guild_id), int(channel_id)) + return {"status": "DELETED", "message_id": message_id} + + @parse_request + async def put_required_roles(self, + guild_id: int, + role_ids: List[int] = None + ): + """ + Set or edit RFR required roles for a guild + + :param guild_id: Guild ID of RFR message + :param role_ids: List of required role IDs + :return: + """ + core.edit_guild_rfr_required_roles(self._bot, guild_id, role_ids) + return core.rfr_list_guild_required_roles(self._bot.get_guild(int(guild_id))) + + @parse_request + async def get_required_roles(self, guild_id: int): + """ + Get RFR required roles for a guild + + :param guild_id: Guild ID of RFR message + :return: + """ + return core.rfr_list_guild_required_roles(self._bot.get_guild(int(guild_id))) + + +def setup(bot: Bot): + """ + Load this cog to the KoalaBot. + :param bot: the bot client for KoalaBot + """ + sub_app = web.Application() + endpoint = RfrEndpoint(bot) + endpoint.register(sub_app) + getattr(bot, "koala_web_app").add_subapp('/{}'.format(RFR_ENDPOINT), sub_app) + logger.info("RFR API is ready.") diff --git a/koala/cogs/react_for_role/cog.py b/koala/cogs/react_for_role/cog.py index 2d0f797e..ea8340f9 100644 --- a/koala/cogs/react_for_role/cog.py +++ b/koala/cogs/react_for_role/cog.py @@ -18,13 +18,13 @@ from discord.ext import commands # Own modules +from . import core import koalabot from koala.colours import KOALA_GREEN from koala.utils import wait_for_message from koala.db import insert_extension -from .db import ReactForRoleDBManager +from .db import get_rfr_message, get_rfr_message_emoji_roles, get_guild_rfr_messages, get_guild_rfr_roles, get_guild_rfr_required_roles from .log import logger -from .utils import CUSTOM_EMOJI_REGEXP, UNICODE_EMOJI_REGEXP def rfr_is_enabled(ctx): @@ -47,10 +47,9 @@ class ReactForRole(commands.Cog): A discord.py cog pertaining to a React for Role system to allow for automation in getting roles. """ - def __init__(self, bot: discord.Client): + def __init__(self, bot): self.bot = bot insert_extension("ReactForRole", 0, True, True) - self.rfr_database_manager = ReactForRoleDBManager() @commands.check(koalabot.is_guild_channel) @commands.check(koalabot.is_admin) @@ -181,15 +180,13 @@ async def rfr_create_message(self, ctx: commands.Context): desc: str = msg.content await ctx.send(f"Okay, the description of the message will be \"{desc}\".\n Okay, " f"I'll create the react for role message now.") - embed: discord.Embed = discord.Embed(title=title, description=desc, colour=KOALA_GREEN) - embed.set_footer(text="ReactForRole") - embed.set_thumbnail( - url="https://cdn.discordapp.com/attachments/737280260541907015/752024535985029240/discord1.png") - rfr_msg: discord.Message = await channel.send(embed=embed) - self.rfr_database_manager.add_rfr_message(ctx.guild.id, channel.id, rfr_msg.id) + + rfr_msg_id = (await core.create_rfr_message(self.bot, ctx.guild.id, channel.id, title, desc, KOALA_GREEN)).message_id + # TODO - Get this working, for some reason we get 403 currently + # await core.setup_rfr_reaction_permissions(ctx.guild, channel, self.bot) await self.overwrite_channel_add_reaction_perms(ctx.guild, channel) await ctx.send( - f"Your react for role message ID is {rfr_msg.id}, it's in {channel.mention}. You can use the other " + f"Your react for role message ID is {rfr_msg_id}, it's in {channel.mention}. You can use the other " "k!rfr subcommands to change the message and add functionality as required.") await del_msg.delete() @@ -209,10 +206,7 @@ async def rfr_delete_message(self, ctx: commands.Context): await ctx.send("Please confirm that you would indeed like to delete the react for role message.") if (await self.prompt_for_input(ctx, "Y/N")).lstrip().strip().upper() == "Y": await ctx.send("Ok") - rfr_msg_row = self.rfr_database_manager.get_rfr_message(ctx.guild.id, channel.id, msg.id) - self.rfr_database_manager.remove_rfr_message_emoji_roles(rfr_msg_row[3]) - self.rfr_database_manager.remove_rfr_message(ctx.guild.id, channel.id, msg.id) - await msg.delete() + await core.delete_rfr_message(self.bot, msg.id, ctx.guild.id, channel.id) await ctx.send("ReactForRole Message deleted") else: await ctx.send("Cancelled command.") @@ -234,14 +228,13 @@ async def rfr_edit_description(self, ctx: commands.Context): await ctx.send("Okay, this will edit the description of an existing react for role message. I'll need some " "details first though.") msg, channel = await self.get_rfr_message_from_prompts(ctx) - embed = self.get_embed_from_message(msg) + embed = core.get_embed_from_message(msg) await ctx.send(f"Your current description is {embed.description}. Please enter your new description.") desc = await self.prompt_for_input(ctx, "description") if desc != "": await ctx.send(f"Your new description would be {desc}. Please confirm that you'd like this change.") if (await self.prompt_for_input(ctx, "Y/N")).lstrip().strip().upper() == "Y": - embed.description = desc - await msg.edit(embed=embed) + await core.rfr_edit(msg, description=desc) else: await ctx.send("Okay, cancelling command.") else: @@ -260,14 +253,13 @@ async def rfr_edit_title(self, ctx: commands.Context): await ctx.send("Okay, this will edit the title of an existing react for role message. I'll need some details " "first though.") msg, channel = await self.get_rfr_message_from_prompts(ctx) - embed = self.get_embed_from_message(msg) + embed = core.get_embed_from_message(msg) await ctx.send(f"Your current title is {embed.title}. Please enter your new title.") title = await self.prompt_for_input(ctx, "title") if title != "": await ctx.send(f"Your new title would be {title}. Please confirm that you'd like this change.") if (await self.prompt_for_input(ctx, "Y/N")).lstrip().strip().upper() == "Y": - embed.title = title - await msg.edit(embed=embed) + await core.rfr_edit(msg, title=title) else: await ctx.send("Okay, cancelling command.") else: @@ -286,7 +278,7 @@ async def rfr_edit_thumbnail(self, ctx: commands.Context): await ctx.send("Okay, this will edit the thumbnail of a react for role message. I'll need some details first " "though.") msg, channel = await self.get_rfr_message_from_prompts(ctx) - embed = self.get_embed_from_message(msg) + embed = core.get_embed_from_message(msg) if not embed: logger.error( f"RFR: Can't find embed for message id {msg.id}, channel {channel.id}, guild id {ctx.guild.id}.") @@ -300,15 +292,13 @@ async def rfr_edit_thumbnail(self, ctx: commands.Context): logger.error(f"Attachment url not found, details : {image}") raise commands.BadArgument("Couldn't get an image from the message you sent.") else: - embed.set_thumbnail(url=str(image.url)) - await msg.edit(embed=embed) + await core.rfr_edit(msg, thumbnail_url=str(image.url)) await ctx.send("Okay, set the thumbnail of the thumbnail to your desired image. This will error if you " "delete the message you sent with the image, so make sure you don't.") elif isinstance(image, str): # no attachment in message, just a raw URL in content img_url = await self.get_image_from_url(ctx, image) - embed.set_thumbnail(url=img_url) - await msg.edit(embed=embed) + await core.rfr_edit(msg, thumbnail_url=img_url) await ctx.send("Okay, set the thumbnail of the thumbnail to your desired image.") else: raise commands.BadArgument("Couldn't get an image from the message you sent.") @@ -351,25 +341,13 @@ async def rfr_edit_inline(self, ctx: commands.Context): await ctx.send( "Keep in mind that this process may take a while if you have a lot of RFR messages on your " "server.") - # fetch rfr messages - guild: discord.Guild = ctx.guild - text_channels: List[discord.TextChannel] = guild.text_channels - guild_rfr_messages = self.rfr_database_manager.get_guild_rfr_messages(guild.id) - for rfr_message in guild_rfr_messages: - channel: discord.TextChannel = discord.utils.get(text_channels, id=rfr_message[1]) - msg: discord.Message = await channel.fetch_message(rfr_message[2]) - embed: discord.Embed = self.get_embed_from_message(msg) - length = self.get_number_of_embed_fields(embed) - for i in range(length): - field = embed.fields[i] - embed.set_field_at(i, name=field.name, value=field.value, inline=change_all == "Y") - await msg.edit(embed=embed) + await core.use_inline_rfr_all(ctx.guild) await ctx.send("Okay, the process should be finished now. Please check.") elif input_comm.lstrip().rstrip().lower() == "specific": # try and get specific message await ctx.send("Okay, I'll need the information about the specific rfr message.") msg, channel = await self.get_rfr_message_from_prompts(ctx) - embed: discord.Embed = self.get_embed_from_message(msg) + embed: discord.Embed = core.get_embed_from_message(msg) if not embed: await ctx.send("Couldn't get embed, is this an RFR message?") else: @@ -383,11 +361,7 @@ async def rfr_edit_inline(self, ctx: commands.Context): await ctx.send("Invalid input, cancelling command") else: await ctx.send("Okay, I'll change it as requested.") - length = self.get_number_of_embed_fields(embed) - for i in range(length): - field = embed.fields[i] - embed.set_field_at(i, name=field.name, value=field.value, inline=yes_no == "Y") - await msg.edit(embed=embed) + await core.use_inline_rfr_specific(msg) await ctx.send("Okay, should be done. Please check.") @commands.check(koalabot.is_admin) @@ -398,19 +372,19 @@ async def rfr_fix_embed(self, ctx: commands.Context): Cosmetic fix method if the bot ever has a moment and doesn't react with the correct emojis/has duplicates. """ msg, chnl = await self.get_rfr_message_from_prompts(ctx) - await self.overwrite_channel_add_reaction_perms(chnl.guild, chnl) - emb = self.get_embed_from_message(msg) + await core.setup_rfr_reaction_permissions(chnl.guild, chnl, self.bot) + emb = core.get_embed_from_message(msg) reacts: List[Union[discord.PartialEmoji, discord.Emoji, str]] = [x.emoji for x in msg.reactions] if not emb: logger.error( f"RFR: Can't find embed for message id {msg.id}, channel {chnl.id}, guild id {ctx.guild.id}.") else: - er_id, _, _, _ = self.rfr_database_manager.get_rfr_message(ctx.guild.id, chnl.id, msg.id) + er_id, _, _, _ = get_rfr_message(ctx.guild.id, chnl.id, msg.id) if not er_id: logger.error( f"RFR: Can't find rfr message with {msg.id}, channel {chnl.id}, guild id {ctx.guild.id}. DB ER_ID : {er_id}") else: - rfr_er = self.rfr_database_manager.get_rfr_message_emoji_roles(er_id) + rfr_er = get_rfr_message_emoji_roles(er_id) if not rfr_er: logger.error( f"RFR: Can't retrieve RFR message (ER_ID: {er_id})'s emoji role combinations.") @@ -453,12 +427,12 @@ async def rfr_add_roles_to_msg(self, ctx: commands.Context): "Okay. This will add roles to an already created react for role message. I'll need some details first " "though.") msg, channel = await self.get_rfr_message_from_prompts(ctx) - rfr_msg_row = self.rfr_database_manager.get_rfr_message(ctx.guild.id, channel.id, msg.id) + rfr_msg_row = get_rfr_message(ctx.guild.id, channel.id, msg.id) if not rfr_msg_row: raise commands.CommandError("Message ID given is not that of a react for role message.") await ctx.send("Okay, found the message you want to add to.") - remaining_slots = 20 - self.get_number_of_embed_fields(self.get_embed_from_message(msg)) + remaining_slots = 20 - core.get_number_of_embed_fields(core.get_embed_from_message(msg)) if remaining_slots == 0: await ctx.send( @@ -469,16 +443,11 @@ async def rfr_add_roles_to_msg(self, ctx: commands.Context): await ctx.send( "Okay, I'll continue then. The new message will have the same title and description as the " "old one.") - old_embed = self.get_embed_from_message(msg) - embed: discord.Embed = discord.Embed(title=old_embed.title, description=old_embed.description) - embed.set_thumbnail( - url=koalabot.KOALA_IMAGE_URL) - msg: discord.Message = await channel.send(embed=embed) - msg_id = msg.id - channel = msg.channel - self.rfr_database_manager.add_rfr_message(ctx.guild.id, channel.id, msg_id) - await ctx.send(f"Okay, the new message has ID {msg.id} and is in {msg.channel.mention}.") - rfr_msg_row = self.rfr_database_manager.get_rfr_message(ctx.guild.id, channel.id, msg_id) + old_embed = core.get_embed_from_message(msg) + rfr_msg_id = (await core.create_rfr_message(self.bot, ctx.guild.id, channel.id, + title=old_embed.title, description=old_embed.description, + colour=KOALA_GREEN)).message_id + await ctx.send(f"Okay, the new message has ID {rfr_msg_id} and is in {msg.channel.mention}.") else: await ctx.send("Okay, I'll stop the command then.") return @@ -492,35 +461,10 @@ async def rfr_add_roles_to_msg(self, ctx: commands.Context): input_role_emojis = (await wait_for_message(self.bot, ctx, 180))[0].content emoji_role_list = await self.parse_emoji_and_role_input_str(ctx, input_role_emojis, remaining_slots) - rfr_embed = self.get_embed_from_message(msg) - - for emoji_role in emoji_role_list: - discord_emoji = emoji_role[0] - role = emoji_role[1] - - if discord_emoji in [x.name for x in rfr_embed.fields]: - await ctx.send("Found duplicate emoji in the message, I'm not accepting it.") - elif role in [x.value for x in rfr_embed.fields]: - await ctx.send("Found duplicate role in the message, I'm not accepting it.") - else: - if isinstance(discord_emoji, str): - self.rfr_database_manager.add_rfr_message_emoji_role(rfr_msg_row[3], emoji.demojize(discord_emoji), - role.id) - else: - self.rfr_database_manager.add_rfr_message_emoji_role(rfr_msg_row[3], str(discord_emoji), role.id) - rfr_embed.add_field(name=str(discord_emoji), value=role.mention, inline=False) - await msg.add_reaction(discord_emoji) - - if isinstance(discord_emoji, str): - logger.info( - f"ReactForRole: Added role ID {str(role.id)} to rfr message (channel, guild) {msg.id} " - f"({str(channel.id)}, {str(ctx.guild.id)}) with emoji {discord_emoji}.") - else: - logger.info( - f"ReactForRole: Added role ID {str(role.id)} to rfr message (channel, guild) {msg.id} " - f"({str(channel.id)}, {str(ctx.guild.id)}) with emoji {discord_emoji.id}.") - - await msg.edit(embed=rfr_embed) + duplicateRolesFound, duplicateEmojisFound, edited_msg = core.rfr_add_emoji_role(ctx.guild, channel, + msg, emoji_role_list) + if (duplicateEmojisFound): await ctx.send("Found duplicate emoji in the message, I'm not accepting it.") + if (duplicateRolesFound): await ctx.send("Found duplicate roles in the message, I'm not accepting it.") await ctx.send("Okay, you should see the message with its new emojis now.") @commands.check(koalabot.is_admin) @@ -542,12 +486,12 @@ async def rfr_remove_roles_from_msg(self, ctx: commands.Context): "Okay, this will remove roles from an already existing react for role message. I'll need some details first" " though.") msg, channel = await self.get_rfr_message_from_prompts(ctx) - rfr_msg_row = self.rfr_database_manager.get_rfr_message(ctx.guild.id, channel.id, msg.id) + rfr_msg_row = get_rfr_message(ctx.guild.id, channel.id, msg.id) if not rfr_msg_row: raise commands.CommandError("Message ID given is not that of a react for role message.") await ctx.send("Okay, found the message you want to remove roles from.") - remaining_slots = self.get_number_of_embed_fields(self.get_embed_from_message(msg)) + remaining_slots = core.get_number_of_embed_fields(core.get_embed_from_message(msg)) if remaining_slots == 0: await ctx.send( @@ -556,8 +500,7 @@ async def rfr_remove_roles_from_msg(self, ctx: commands.Context): if (await self.prompt_for_input(ctx, "Y/N")).lstrip().strip().upper() == "Y": await ctx.send("Okay, deleting that message and removing it from the database.") - self.rfr_database_manager.remove_rfr_message(ctx.guild.id, channel.id, msg.id) - await msg.delete() + await core.delete_rfr_message(self.bot, msg.id, ctx.guild.id, channel.id) await ctx.send("Okay, deleted that react for role message. Have a nice day.") return else: @@ -571,54 +514,20 @@ async def rfr_remove_roles_from_msg(self, ctx: commands.Context): input_emoji_roles = (await wait_for_message(self.bot, ctx, 120))[0].content wanted_removals = await self.parse_emoji_or_roles_input_str(ctx, input_emoji_roles) - rfr_embed: discord.Embed = self.get_embed_from_message(msg) - rfr_embed_fields = rfr_embed.fields - new_embed = discord.Embed(title=rfr_embed.title, description=rfr_embed.description, - colour=KOALA_GREEN) - new_embed.set_thumbnail( - url="https://cdn.discordapp.com/attachments/737280260541907015/752024535985029240/discord1.png") - new_embed.set_footer(text="ReactForRole") - removed_field_indexes = [] - reactions_to_remove: List[discord.Reaction] = [] - - for row in wanted_removals: - if isinstance(row, discord.Emoji) or isinstance(row, str): - field_index = [x.name for x in rfr_embed_fields].index(str(row)) - if isinstance(row, str): - self.rfr_database_manager.remove_rfr_message_emoji_role(rfr_msg_row[3], - emoji_raw=emoji.demojize(row)) - else: - self.rfr_database_manager.remove_rfr_message_emoji_role(rfr_msg_row[3], emoji_raw=row) - else: - # row is instance of role - field_index = [x.value for x in rfr_embed_fields].index(row.mention) - self.rfr_database_manager.remove_rfr_message_emoji_role(rfr_msg_row[3], role_id=row.id) - field = rfr_embed_fields[field_index] - removed_field_indexes.append(field_index) - reaction_emoji = await self.get_first_emoji_from_str(ctx, field.name) - reaction: discord.Reaction = [x for x in msg.reactions if str(x.emoji) == str(reaction_emoji)][0] - reactions_to_remove.append(reaction) + new_embed, errors = core.rfr_remove_emojis_roles(self.bot, ctx.guild, msg, rfr_msg_row, wanted_removals) + for e in errors: + await ctx.send(e) - new_embed_fields = [field for field in rfr_embed_fields if - rfr_embed_fields.index(field) not in removed_field_indexes] - - for field in new_embed_fields: - new_embed.add_field(name=field.name, value=field.value, inline=False) - - if self.get_number_of_embed_fields(new_embed) == 0: + if core.get_number_of_embed_fields(new_embed) == 0: await ctx.send("I see you've removed all emoji-role combinations from this react for role message. " "Would you like to delete this message?") if (await self.prompt_for_input(ctx, "Y/N")).lstrip().strip().upper() == "Y": await ctx.send("Okay, I'll delete the message then.") - self.rfr_database_manager.remove_rfr_message(ctx.guild.id, channel.id, msg.id) - await msg.delete() + await core.delete_rfr_message(self.bot, msg.id, ctx.guild.id, channel.id) return - for reaction in reactions_to_remove: - await reaction.clear() - await msg.edit(embed=new_embed) await ctx.send("Okay, I've removed those options from the react for role message.") @commands.Cog.listener() @@ -632,7 +541,7 @@ async def on_raw_reaction_add(self, payload: discord.RawReactionActionEvent): """ if payload.guild_id is not None: if not payload.member.bot: - rfr_message = self.rfr_database_manager.get_rfr_message(payload.guild_id, payload.channel_id, + rfr_message = get_rfr_message(payload.guild_id, payload.channel_id, payload.message_id) if not rfr_message: return @@ -651,7 +560,7 @@ async def on_raw_reaction_add(self, payload: discord.RawReactionActionEvent): await member_role[0].add_roles(member_role[1]) else: # Remove all rfr roles from member - role_ids = self.rfr_database_manager.get_guild_rfr_roles(payload.guild_id) + role_ids = get_guild_rfr_roles(payload.guild_id) roles: List[discord.Role] = [] for role_id in role_ids: role = discord.utils.get(member_role[0].guild.roles, id=role_id) @@ -661,7 +570,7 @@ async def on_raw_reaction_add(self, payload: discord.RawReactionActionEvent): for role_to_remove in roles: await member_role[0].remove_roles(role_to_remove) # Remove members' reaction from all rfr messages in guild - guild_rfr_messages = self.rfr_database_manager.get_guild_rfr_messages(payload.guild_id) + guild_rfr_messages = get_guild_rfr_messages(payload.guild_id) if not guild_rfr_messages: logger.error( f"ReactForRole: Guild RFR messages is empty on raw reaction add. Please check" @@ -678,39 +587,37 @@ async def on_raw_reaction_add(self, payload: discord.RawReactionActionEvent): @commands.check(koalabot.is_admin) @commands.check(rfr_is_enabled) @react_for_role_group.command("addRequiredRole") - async def rfr_add_guild_required_role(self, ctx: commands.Context, role_str: str): + async def rfr_add_guild_required_role(self, ctx: commands.Context, role: discord.Role): """ Adds a role to perms to use rfr functionality in a server, so you can specify that you need, e.g. "@Student" to be able to use rfr functionality in the server. It's server-wide permissions handling however. By default anyone can use rfr functionality in the server. User needs to have admin perms to use. :param ctx: Context of the command - :param role_str: Role ID/name/mention + :param role: Role ID/name/mention :return: """ try: - role: discord.Role = await commands.RoleConverter().convert(ctx, role_str) + core.add_guild_rfr_required_role(ctx.guild, role.id) await ctx.send(f"Okay, I'll add {role.name} to the list of roles required for RFR usage on the server.") - self.rfr_database_manager.add_guild_rfr_required_role(ctx.guild.id, role.id) except (commands.CommandError, commands.BadArgument): await ctx.send("Found an issue with your provided argument, couldn't get an actual role. Please try again.") @commands.check(koalabot.is_admin) @commands.check(rfr_is_enabled) @react_for_role_group.command("removeRequiredRole") - async def rfr_remove_guild_required_role(self, ctx: commands.Context, role_str: str): + async def rfr_remove_guild_required_role(self, ctx: commands.Context, role: discord.Role): """ Removes a role from perms for use of rfr functionality in a server, so you can specify that you need, e.g. "@Student" to be able to use rfr functionality in the server. It's server-wide permissions handling however. By default anyone can use rfr functionality in the server. User needs to have admin perms to use. :param ctx: Context of the command - :param role_str: Role ID/name/mention + :param role: Role ID/name/mention :return: """ try: - role: discord.Role = await commands.RoleConverter().convert(ctx, role_str) + core.remove_guild_rfr_required_role(ctx.guild, role.id) await ctx.send( f"Okay, I'll remove {role.name} from the list of roles required for RFR usage on the server.") - self.rfr_database_manager.remove_guild_rfr_required_role(ctx.guild.id, role.id) except (commands.CommandError, commands.BadArgument): await ctx.send("Found an issue with your provided argument, couldn't get an actual role. Please try again.") @@ -724,7 +631,7 @@ async def rfr_list_guild_required_roles(self, ctx: commands.Context): :param ctx: Context of the command. :return: """ - role_ids = self.rfr_database_manager.get_guild_rfr_required_roles(ctx.guild.id) + role_ids = core.rfr_list_guild_required_roles(ctx.guild).role_ids msg_str = "You will need one of these roles to react to rfr messages on this server:\n" for role_id in role_ids: @@ -749,7 +656,7 @@ async def on_raw_reaction_remove(self, payload: discord.RawReactionActionEvent): """ if payload.guild_id is not None: - rfr_message = self.rfr_database_manager.get_rfr_message(payload.guild_id, payload.channel_id, + rfr_message = get_rfr_message(payload.guild_id, payload.channel_id, payload.message_id) if not rfr_message: return @@ -766,7 +673,7 @@ def can_have_rfr_role(self, member: discord.Member) -> bool: :param member: Member to check rfr perms for :return: True if member has one of the required roles, or if there are no required roles. False otherwise """ - required_roles: List[int] = self.rfr_database_manager.get_guild_rfr_required_roles(member.guild.id) + required_roles: List[int] = get_guild_rfr_required_roles(member.guild.id) if not required_roles or len(required_roles) == 0: return True return any(x in required_roles for x in [y.id for y in member.roles]) @@ -789,7 +696,7 @@ async def get_rfr_message_from_prompts(self, ctx: commands.Context) -> Tuple[dis msg = await channel.fetch_message(msg_id) if not msg: raise commands.CommandError("Invalid Message ID given.") - rfr_msg_row = self.rfr_database_manager.get_rfr_message(ctx.guild.id, channel.id, msg_id) + rfr_msg_row = get_rfr_message(ctx.guild.id, channel.id, msg_id) if not rfr_msg_row: raise commands.CommandError("Message ID given is not that of a react for role message.") return msg, channel @@ -820,7 +727,7 @@ async def get_role_member_info(self, emoji_reacted: discord.PartialEmoji, guild_ message: discord.Message = await channel.fetch_message(message_id) if not message: return - embed: discord.Embed = self.get_embed_from_message(message) + embed: discord.Embed = core.get_embed_from_message(message) if emoji_reacted.is_unicode_emoji(): rep = emoji.emojize(emoji_reacted.name) @@ -839,6 +746,9 @@ async def get_role_member_info(self, emoji_reacted: discord.PartialEmoji, guild_ elif emoji_reacted.is_custom_emoji(): rep = str(emoji_reacted) field = await self.get_field_by_emoji(embed, rep) + if not field: + # Look for animated version + field = await self.get_field_by_emoji(embed, rep[0]+"a"+rep[1:]) if not field: return role_str = field @@ -863,15 +773,21 @@ async def parse_emoji_and_role_input_str(self, ctx: commands.Context, input_str: :return: List of Emoji-Role pairs parsed from the input message. """ rows = input_str.splitlines() + arr = [] for row in rows: emoji_role = row.split(',') + # print(emoji_role) + if (len(emoji_role) < 2): + continue if len(emoji_role) > 2: - raise commands.BadArgument("Too many categories/etc on one line.") - emoji: Union[discord.Emoji, str] = await self.get_first_emoji_from_str(ctx, emoji_role[0].strip()) + raise commands.BadArgument("Too many/little categories/etc on one line.") + emoji, err = await core.get_first_emoji_from_str(ctx.bot, ctx.guild, emoji_role[0].strip()) + if not emoji: - await ctx.send(f"Yeah, didn't find emoji for `{emoji_role[0]}`") + await ctx.send(f"Yeah, didn't find emoji for `{emoji_role[0]}` - {err}") continue + role = await commands.RoleConverter().convert(ctx, emoji_role[1].lstrip().rstrip()) arr.append((emoji, role)) if len(arr) == remaining_slots: @@ -895,7 +811,9 @@ async def parse_emoji_or_roles_input_str(self, ctx: commands.Context, input_str: arr = [] for row in rows: # Try and match it to an raw_emoji first - raw_emoji = await self.get_first_emoji_from_str(ctx, row.strip()) + raw_emoji, err = await core.get_first_emoji_from_str(self.bot, ctx.guild, row.strip()) + if err: + await ctx.send(err) if not raw_emoji: role = await commands.RoleConverter().convert(ctx, row.strip()) if not role: @@ -906,6 +824,7 @@ async def parse_emoji_or_roles_input_str(self, ctx: commands.Context, input_str: arr.append(raw_emoji) return arr + async def prompt_for_input(self, ctx: commands.Context, input_type: str) -> Union[discord.Attachment, str]: """ Prompts a user for input in the form of a message. Has a forced timer of 60 seconds, because it basically just @@ -943,6 +862,7 @@ async def overwrite_channel_add_reaction_perms(self, guild: discord.Guild, chann for bot_member in bot_members: await channel.set_permissions(bot_member, overwrite=overwrite) + async def is_user_alive(self, ctx: commands.Context): """ Prompts user for message to check if they're alive. Any message will do. We hope they're alive anyways. @@ -954,60 +874,6 @@ async def is_user_alive(self, ctx: commands.Context): return False return True - def get_embed_from_message(self, msg: discord.Message) -> Optional[discord.Embed]: - """ - Gets the embed from a given message. Yup. That's it. - :param msg: Message to check - :return: Returns the embed if there is one. If there isn't returns None - """ - if not msg: - return None - try: - embed = msg.embeds[0] - if not embed: - return None - return embed - except IndexError: - return None - - def get_number_of_embed_fields(self, embed: discord.Embed) -> int: - """ - Gets the number of fields in an embed. - :param embed: Embed to check - :return: Number of embed fields. - """ - return len(embed.fields) - - async def get_first_emoji_from_str(self, ctx: commands.Context, content: str) -> Optional[ - Union[discord.Emoji, str]]: - """ - Gets the first emoji in a string input, custom or not. Doesn't work with custom emojis the bot doesn't have - access to. - :param ctx: Context of the original command - :param content: Message content - :return: Emoji if there is a valid one. Otherwise None. - """ - # First check for a custom discord emoji in the string - search_result = CUSTOM_EMOJI_REGEXP.search(content) - if not search_result: - # Check for a unicode emoji in the string - search_result = UNICODE_EMOJI_REGEXP.search(content) - if not search_result: - return None - return content - else: - emoji_str = search_result.group().strip() - try: - discord_emoji: discord.Emoji = await commands.EmojiConverter().convert(ctx, emoji_str) - return discord_emoji - except commands.CommandError: - await ctx.send( - "An error occurred when trying to get the emoji. Please contact the bot developers for support.") - return None - except commands.BadArgument: - await ctx.send("Couldn't get the emoji you used - is it from this server or a server I'm in?") - return None - async def get_field_by_emoji(self, embed: discord.Embed, emoji: Optional[str]): """ Get the specific field value of an rfr embed by the string representation of the emoji in the field name diff --git a/koala/cogs/react_for_role/core.py b/koala/cogs/react_for_role/core.py new file mode 100644 index 00000000..fcd58dc6 --- /dev/null +++ b/koala/cogs/react_for_role/core.py @@ -0,0 +1,375 @@ +from ast import Tuple +from typing import * + +import discord +from discord.ext.commands import Bot +from discord.ext import commands +import emoji + +import koalabot +from . import db +from .db import get_rfr_message +from .dto import ReactMessage, ReactRole, RequiredRoles +from .log import logger + +from koala.db import assign_session +from koala.colours import KOALA_GREEN +from .utils import CUSTOM_EMOJI_REGEXP, UNICODE_EMOJI_REGEXP + +# Constants + +koala_logo = "https://cdn.discordapp.com/attachments/737280260541907015/752024535985029240/discord1.png" + + +# Variables +# current_activity = None + +def create_ctx(bot: Bot, guild: discord.Guild): + return {'bot': bot, 'guild': guild} + + +@assign_session +async def get_rfr_message_dto(bot: koalabot.KoalaBot, message_id: int, guild_id: int, channel_id: int, + **kwargs): + guild = bot.get_guild(guild_id) + channel = guild.get_channel(channel_id) + message: discord.Message = await channel.fetch_message(message_id) + + rfr_embed = get_embed_from_message(message) + + _, _, _, emoji_role_id = db.get_rfr_message(guild_id, channel_id, message_id, **kwargs) + roles_list = db.get_rfr_message_emoji_roles(emoji_role_id, **kwargs) + + return ReactMessage( + message_id=message_id, + guild_id=guild_id, + channel_id=channel_id, + title=rfr_embed.title, + description=rfr_embed.description, + thumbnail=rfr_embed.thumbnail.url, + colour=rfr_embed.colour.__str__(), + inline=len(rfr_embed.fields) > 0 and rfr_embed.fields[0].inline, + roles=[ReactRole(role[1], role[2]) for role in roles_list] + ) + + +@assign_session +async def create_rfr_message(bot: koalabot.KoalaBot, guild_id: int, channel_id: int, title: str, description: str, + colour: discord.Colour, thumbnail: str = None, inline: bool = None, + roles: List[Tuple[Union[discord.Emoji, str], discord.Role]] = None, + **kwargs) -> ReactMessage: + guild = bot.get_guild(guild_id) + channel = guild.get_channel(channel_id) + + embed: discord.Embed = discord.Embed(title=title, description=description, colour=colour) + embed.set_footer(text="ReactForRole") + if thumbnail is None: + embed.set_thumbnail(url=koala_logo) + else: + embed.set_thumbnail(url=thumbnail) + + rfr_msg: discord.Message = await channel.send(embed=embed) + db.add_rfr_message(guild_id, channel_id, rfr_msg.id, **kwargs) + + if roles is not None: + await rfr_add_emoji_role(guild, channel, rfr_msg, roles, **kwargs) + + if inline: + await use_inline_rfr_specific(rfr_msg) + + return await get_rfr_message_dto(bot, rfr_msg.id, guild_id, channel_id, **kwargs) + + +@assign_session +async def update_rfr_message(bot: koalabot.KoalaBot, message_id: int, guild_id: int, channel_id: int, + title: str, description: str, colour: discord.Colour, + thumbnail: str, inline: bool, + roles: List[Tuple[Union[discord.Emoji, str], discord.Role]], + **kwargs): + guild = bot.get_guild(guild_id) + channel = guild.get_channel(channel_id) + + if roles is not None: + await rfr_edit_emoji_role(bot, message_id, guild_id, channel_id, roles, **kwargs) + + await rfr_edit(await channel.fetch_message(message_id), title=title, description=description, thumbnail_url=thumbnail, colour=colour) + + if inline is not None: + await use_inline_rfr_specific(await channel.fetch_message(message_id)) + + return await get_rfr_message_dto(bot, message_id, guild_id, channel_id, **kwargs) + + +@assign_session +async def delete_rfr_message(bot: koalabot.KoalaBot, message_id: int, guild_id: int, channel_id: int, **kwargs): + rfr_msg_row = db.get_rfr_message(guild_id, channel_id, message_id, **kwargs) + db.remove_rfr_message_emoji_roles(rfr_msg_row[3], **kwargs) + db.remove_rfr_message(guild_id, channel_id, message_id, **kwargs) + + guild = bot.get_guild(guild_id) + channel = guild.get_channel(channel_id) + message = await channel.fetch_message(message_id) + + await message.delete() + + +@assign_session +async def use_inline_rfr_all(guild: discord.Guild, **kwargs): + text_channels: List[discord.TextChannel] = guild.text_channels + guild_rfr_messages = db.get_guild_rfr_messages(guild.id, **kwargs) + for rfr_message in guild_rfr_messages: + channel: discord.TextChannel = discord.utils.get(text_channels, id=rfr_message[1]) + msg: discord.Message = await channel.fetch_message(rfr_message[2]) + embed: discord.Embed = get_embed_from_message(msg) + length = get_number_of_embed_fields(embed) + for i in range(length): + field = embed.fields[i] + embed.set_field_at(i, name=field.name, value=field.value, inline=True) + await msg.edit(embed=embed) + + +async def use_inline_rfr_specific(msg: discord.Message): + rfr_embed = get_embed_from_message(msg) + length = get_number_of_embed_fields(rfr_embed) + for i in range(length): + field = rfr_embed.fields[i] + rfr_embed.set_field_at(i, name=field.name, value=field.value, inline=True) + await msg.edit(embed=rfr_embed) + + +async def rfr_edit(message: discord.Message, *, + title: str = None, description: str = None, thumbnail_url: str = None, colour: discord.Colour = None): + embed = get_embed_from_message(message) + if title is not None: + embed.title = title + if description is not None: + embed.description = description + if thumbnail_url is not None: + embed.set_thumbnail(url=thumbnail_url) + if colour is not None: + embed.colour = colour + return await message.edit(embed=embed) + + +@assign_session +async def rfr_remove_emojis_roles(bot: Bot, guild: discord.Guild, msg: discord.Message, + rfr_msg_row: Tuple[int, int, int, int], + wanted_removals: List[Union[discord.Emoji, str, discord.Role]], **kwargs): + rfr_embed: discord.Embed = get_embed_from_message(msg) + rfr_embed_fields = rfr_embed.fields + new_embed = rfr_embed.copy() + new_embed.clear_fields() + removed_field_indexes = [] + reactions_to_remove: List[discord.Reaction] = [] + errors = [] + + for row in wanted_removals: + if isinstance(row, discord.Emoji) or isinstance(row, str): + field_index = [x.name for x in rfr_embed_fields].index(str(row)) + if isinstance(row, str): + db.remove_rfr_message_emoji_role(rfr_msg_row[3], emoji_raw=emoji.demojize(row), **kwargs) + else: + db.remove_rfr_message_emoji_role(rfr_msg_row[3], emoji_raw=row, **kwargs) + else: + # row is instance of role + field_index = [x.value for x in rfr_embed_fields].index(row.mention) + db.remove_rfr_message_emoji_role(rfr_msg_row[3], role_id=row.id, **kwargs) + + field = rfr_embed_fields[field_index] + removed_field_indexes.append(field_index) + reaction_emoji, err = await get_first_emoji_from_str(bot, guild, field.name) + if (err != None): + errors.append(err) + reaction: discord.Reaction = [x for x in msg.reactions if str(x.emoji) == str(reaction_emoji)][0] + reactions_to_remove.append(reaction) + + new_embed_fields = [field for field in rfr_embed_fields if + rfr_embed_fields.index(field) not in removed_field_indexes] + + for field in new_embed_fields: + new_embed.add_field(name=field.name, value=field.value, inline=False) + + for reaction in reactions_to_remove: + await reaction.clear() + await msg.edit(embed=new_embed) + + return new_embed, errors + + +@assign_session +async def rfr_edit_emoji_role(bot: koalabot.KoalaBot, message_id: int, guild_id: int, channel_id: int, + emoji_role_map: List[Tuple[Union[discord.Emoji, str], discord.Role]], + **kwargs): + guild = bot.get_guild(guild_id) + channel = guild.get_channel(channel_id) + + _, _, _, emoji_role_id = db.get_rfr_message(guild_id, channel_id, message_id, **kwargs) + emoji_roles = db.get_rfr_message_emoji_roles(emoji_role_id, **kwargs) + remove_role_map = {emoji.emojize(r[1]): guild.get_role(r[2]) for r in emoji_roles} + add_role_map = {} + + for emoji_str, role in emoji_role_map: + if emoji.emojize(emoji_str) in remove_role_map.keys(): + remove_role_map.pop(emoji.emojize(emoji_str)) + else: + add_role_map[emoji.emojize(emoji_str)] = role + + remove_role_map = [(r, remove_role_map.get(r)) for r in remove_role_map.keys()] + add_role_map = [(r, add_role_map.get(r)) for r in add_role_map.keys()] + + if remove_role_map: + await rfr_remove_emojis_roles(bot, guild, await channel.fetch_message(message_id), get_rfr_message(guild_id, channel_id, message_id, **kwargs), + [r[1] for r in remove_role_map], **kwargs) + + if add_role_map: + await rfr_add_emoji_role(guild, channel, await channel.fetch_message(message_id), add_role_map, **kwargs) + + +@assign_session +async def rfr_add_emoji_role(guild: discord.Guild, channel: discord.TextChannel, + msg: discord.Message, emoji_role_map: List[Tuple[Union[discord.Emoji, str], discord.Role]], + **kwargs): + rfr_embed = get_embed_from_message(msg) + duplicate_roles_found = False + duplicate_emojis_found = False + rfr_msg_row = db.get_rfr_message(guild.id, channel.id, msg.id) + + for emoji_role in emoji_role_map: + discord_emoji = emoji_role[0] + role = emoji_role[1] + + if discord_emoji in [x.name for x in rfr_embed.fields]: + duplicate_emojis_found = True + elif role in [x.value for x in rfr_embed.fields]: + duplicate_roles_found = True + else: + if isinstance(discord_emoji, str): + db.add_rfr_message_emoji_role(rfr_msg_row[3], emoji.demojize(discord_emoji), + role.id, **kwargs) + else: + db.add_rfr_message_emoji_role(rfr_msg_row[3], str(discord_emoji), role.id, **kwargs) + rfr_embed.add_field(name=str(discord_emoji), value=role.mention, inline=False) + await msg.add_reaction(discord_emoji) + + if isinstance(discord_emoji, str): + logger.info( + f"ReactForRole: Added role ID {str(role.id)} to rfr message (channel, guild) {msg.id} " + f"({str(channel.id)}, {str(guild.id)}) with emoji {emoji.demojize(discord_emoji)}.") + else: + logger.info( + f"ReactForRole: Added role ID {str(role.id)} to rfr message (channel, guild) {msg.id} " + f"({str(channel.id)}, {str(guild.id)}) with emoji {discord_emoji.id}.") + + edited_msg = await msg.edit(embed=rfr_embed) + return duplicate_roles_found, duplicate_emojis_found, edited_msg + + +@assign_session +def edit_guild_rfr_required_roles(bot: koalabot.KoalaBot, guild_id: int, role_ids: List[int], **kwargs): + guild = bot.get_guild(guild_id) + + add_role_ids = [] + remove_role_ids: List = rfr_list_guild_required_roles(guild, **kwargs).role_ids + + for role_id in role_ids: + if role_id in remove_role_ids: + remove_role_ids.remove(role_id) + else: + add_role_ids.append(role_id) + + for role_id in remove_role_ids: + remove_guild_rfr_required_role(guild, role_id, **kwargs) + + for role_id in add_role_ids: + add_guild_rfr_required_role(guild, role_id, **kwargs) + + +def add_guild_rfr_required_role(guild: discord.Guild, role_id: int, **kwargs): + db.add_guild_rfr_required_role(guild.id, role_id, **kwargs) + + +def remove_guild_rfr_required_role(guild: discord.Guild, role_id: int, **kwargs): + db.remove_guild_rfr_required_role(guild.id, role_id, **kwargs) + + +def rfr_list_guild_required_roles(guild: discord.Guild, **kwargs): + return RequiredRoles(guild.id, db.get_guild_rfr_required_roles(guild.id, **kwargs)) + + +async def setup_rfr_reaction_permissions(guild: discord.Guild, channel: discord.TextChannel, bot: Bot): + """ + Overwrites a text channel's reaction perms so that nobody can add new reactions to any message sent in the + channel, only the bot, to make sure people don't mess with the system. Relies on roles tending not to be added/ + removed constantly to keep performance satisfactory. + :param guild: Guild that the rfr message is in + :param channel: Channel that the rfr message is in + :return: + """ + # Get the @everyone role. + role: discord.Role = discord.utils.get(guild.roles, id=guild.id) + overwrite: discord.PermissionOverwrite = discord.PermissionOverwrite() + overwrite.update(add_reactions=False) + # TODO - tests fail here with 403, missing 'manage_roles' permission + await channel.set_permissions(role, overwrite=overwrite) + bot_members = [member for member in guild.members if member.bot and member.id == bot.user.id] + overwrite.update(add_reactions=True) + for bot_member in bot_members: + await channel.set_permissions(bot_member, overwrite=overwrite) + + +def get_embed_from_message(msg: discord.Message) -> Optional[discord.Embed]: + """ + Gets the embed from a given message + :param msg: Message to check + :return: Returns the embed if there is one. If there isn't returns None + """ + if not msg: + return None + try: + embed = msg.embeds[0] + if not embed: + return None + return embed + except IndexError: + return None + + +def get_number_of_embed_fields(embed: discord.Embed) -> int: + """ + Gets the number of fields in an embed. + :param embed: Embed to check + :return: Number of embed fields. + """ + return len(embed.fields) + + +async def get_first_emoji_from_str(bot: Bot, guild: discord.Guild, + content: str) -> Tuple[Optional[Union[discord.Emoji, str]], Optional[str]]: + """ + Gets the first emoji in a string input, custom or not. Doesn't work with custom emojis the bot doesn't have + access to. + :param bot: + :param guild: + :param content: Message content + :return: Emoji if there is a valid one. Otherwise None. + """ + + # First check for a custom discord emoji in the string + search_result = CUSTOM_EMOJI_REGEXP.search(str(content)) + if not search_result: + # Check for a unicode emoji in the string + search_result = UNICODE_EMOJI_REGEXP.search(content) + if not search_result: + return None, "No emoji found." + return content, None + else: + emoji_id = int(search_result[:-1].split(":")[-1]) + try: + discord_emoji: discord.Emoji = await guild.fetch_emoji(emoji_id) + if discord_emoji is None: + discord_emoji: discord.Emoji = bot.get_emoji(emoji_id) + return discord_emoji, None + except commands.CommandError: + return None, "An error occurred when trying to get the emoji. Please contact the bot developers for support." + except commands.BadArgument: + return None, "Couldn't get the emoji you used - is it from this server or a server I'm in?" diff --git a/koala/cogs/react_for_role/db.py b/koala/cogs/react_for_role/db.py index df283dc1..46c6aa0a 100644 --- a/koala/cogs/react_for_role/db.py +++ b/koala/cogs/react_for_role/db.py @@ -1,253 +1,256 @@ #!/usr/bin/env python -""" -KoalaBot Reaction Roles Code - -Author: Anan Venkatesh -Commented using reStructuredText (reST) -""" -# Futures - # Built-in/Generic Imports from typing import * import sqlalchemy.exc +import sqlalchemy.orm from sqlalchemy import select, delete, and_ # Own modules from koala.db import session_manager from .log import logger from .models import GuildRFRMessages, RFRMessageEmojiRoles, GuildRFRRequiredRoles +from koala.db import assign_session + + +@assign_session +def add_rfr_message(guild_id: int, channel_id: int, message_id: int, session: sqlalchemy.orm.Session): + """ + Add an rfr message to a guild. Table stores a unique emoji_role_id to prevent the same combination + appearing twice on a given message + :param guild_id: ID of the guild + :param channel_id: ID of the channel the rfr message is in + :param message_id: ID of the rfr message + :return: + """ + session.add( + GuildRFRMessages(guild_id=guild_id, channel_id=channel_id, message_id=message_id)) + session.commit() + + +@assign_session +def add_rfr_message_emoji_role(emoji_role_id: int, emoji_raw: str, role_id: int, session: sqlalchemy.orm.Session): + """ + Add an emoji-role combination to an rfr message. + :param emoji_role_id: unique ID/key + :param emoji_raw: raw emoji representation in string format + :param role_id: ID of the role to give on react + :return: + """ + try: + session.add(RFRMessageEmojiRoles(emoji_role_id=emoji_role_id, emoji_raw=emoji_raw, role_id=role_id)) + session.commit() + except sqlalchemy.exc.IntegrityError: + logger.warning("RFRMessageEmojiRoles already exists for <%s, %s, %s>, continuing", + emoji_role_id, emoji_raw, role_id) + + +@assign_session +def remove_rfr_message_emoji_role(emoji_role_id: int, emoji_raw: str = None, role_id: int = None, + session: sqlalchemy.orm.Session = None): + """ + Remove an emoji-role combination from the rfr message database. Uses the unique emoji_role_id to identify the + specific combo. Only removes one emoji-role combo + :param emoji_role_id: unique ID/key + :param emoji_raw: raw string representation of the emoji + :param role_id: ID of the role to give on react + :return: + """ + if not emoji_raw: + delete_sql = delete(RFRMessageEmojiRoles) \ + .where( + and_( + RFRMessageEmojiRoles.emoji_role_id == emoji_role_id, + RFRMessageEmojiRoles.role_id == role_id + )) + else: + delete_sql = delete(RFRMessageEmojiRoles) \ + .where( + and_( + RFRMessageEmojiRoles.emoji_role_id == emoji_role_id, + RFRMessageEmojiRoles.emoji_raw == emoji_raw + )) + session.execute(delete_sql) + session.commit() + + +@assign_session +def remove_rfr_message_emoji_roles(emoji_role_id: int, session: sqlalchemy.orm.Session): + """ + Removes all emoji-role combos with the same emoji_role_id i.e. on the same message. + :param emoji_role_id: unique ID/key + :return: + """ + delete_sql = delete(RFRMessageEmojiRoles) \ + .where(RFRMessageEmojiRoles.emoji_role_id == emoji_role_id) + + session.execute(delete_sql) + session.commit() + + +@assign_session +def remove_rfr_message(guild_id: int, channel_id: int, message_id: int, session: sqlalchemy.orm.Session): + """ + Removes an rfr message from the rfr message database, and also removes all emoji-role combos as part of it. + :param guild_id: Guild ID of the rfr message + :param channel_id: Channel ID of the rfr message + :param message_id: Message ID of the rfr message + :return: + """ + emoji_role_id = get_rfr_message(guild_id, channel_id, message_id) + if not emoji_role_id: + return + else: + remove_rfr_message_emoji_roles(emoji_role_id[3]) + + delete_sql = delete(GuildRFRMessages) \ + .where(and_(and_( + GuildRFRMessages.guild_id == guild_id, + GuildRFRMessages.channel_id == channel_id), + GuildRFRMessages.message_id == message_id)) + session.execute(delete_sql) + session.commit() + + +@assign_session +def get_rfr_message(guild_id: int, channel_id: int, message_id: int, session: sqlalchemy.orm.Session) -> Optional[ + Tuple[int, int, int, int]]: + """ + Gets the unique rfr message that is specified by the guild ID, channel ID and message ID. + :param guild_id: Guild ID of the rfr message + :param channel_id: Channel ID of the rfr message + :param message_id: Message ID of the rfr message + :return: RFR message info of the specific message if found, otherwise None. + """ + message = session.execute(select(GuildRFRMessages) + .filter_by(guild_id=guild_id, + channel_id=channel_id, + message_id=message_id)).scalars().one_or_none() + if message: + return message.old_format() + else: + return None + + +@assign_session +def get_guild_rfr_messages(guild_id: int, session: sqlalchemy.orm.Session) -> List[Tuple[int, int, int]]: + """ + Gets all rfr messages in a given guild, from the guild ID + :param guild_id: ID of the guild + :return: List of rfr messages in the guild. + """ + messages = session.execute(select(GuildRFRMessages) + .filter_by(guild_id=guild_id)).scalars().all() + return [message.old_format() + for message in messages] -# Libs - -# Constants - - -class ReactForRoleDBManager: - """ - A class for interacting with the KoalaBot ReactForRole database - """ - - def add_rfr_message(self, guild_id: int, channel_id: int, message_id: int): - """ - Add an rfr message to a guild. Table stores a unique emoji_role_id to prevent the same combination - appearing twice on a given message - :param guild_id: ID of the guild - :param channel_id: ID of the channel the rfr message is in - :param message_id: ID of the rfr message - :return: - """ - with session_manager() as session: - session.add( - GuildRFRMessages(guild_id=guild_id, channel_id=channel_id, message_id=message_id)) - session.commit() - - def add_rfr_message_emoji_role(self, emoji_role_id: int, emoji_raw: str, role_id: int): - """ - Add an emoji-role combination to an rfr message. - :param emoji_role_id: unique ID/key - :param emoji_raw: raw emoji representation in string format - :param role_id: ID of the role to give on react - :return: - """ - with session_manager() as session: - try: - session.add(RFRMessageEmojiRoles(emoji_role_id=emoji_role_id, emoji_raw=emoji_raw, role_id=role_id)) - session.commit() - except sqlalchemy.exc.IntegrityError: - logger.warning("RFRMessageEmojiRoles already exists for <%s, %s, %s>, continuing", - emoji_role_id, emoji_raw, role_id) - - def remove_rfr_message_emoji_role(self, emoji_role_id: int, emoji_raw: str = None, role_id: int = None): - """ - Remove an emoji-role combination from the rfr message database. Uses the unique emoji_role_id to identify the - specific combo. Only removes one emoji-role combo - :param emoji_role_id: unique ID/key - :param emoji_raw: raw string representation of the emoji - :param role_id: ID of the role to give on react - :return: - """ - if not emoji_raw: - delete_sql = delete(RFRMessageEmojiRoles)\ - .where( - and_( - RFRMessageEmojiRoles.emoji_role_id == emoji_role_id, - RFRMessageEmojiRoles.role_id == role_id - )) - else: - delete_sql = delete(RFRMessageEmojiRoles)\ - .where( - and_( - RFRMessageEmojiRoles.emoji_role_id == emoji_role_id, - RFRMessageEmojiRoles.emoji_raw == emoji_raw - )) - with session_manager() as session: - session.execute(delete_sql) - session.commit() - - def remove_rfr_message_emoji_roles(self, emoji_role_id: int): - """ - Removes all emoji-role combos with the same emoji_role_id i.e. on the same message. - :param emoji_role_id: unique ID/key - :return: - """ - with session_manager() as session: - delete_sql = delete(RFRMessageEmojiRoles) \ - .where(RFRMessageEmojiRoles.emoji_role_id == emoji_role_id) - - session.execute(delete_sql) - session.commit() - - def remove_rfr_message(self, guild_id: int, channel_id: int, message_id: int): - """ - Removes an rfr message from the rfr message database, and also removes all emoji-role combos as part of it. - :param guild_id: Guild ID of the rfr message - :param channel_id: Channel ID of the rfr message - :param message_id: Message ID of the rfr message - :return: - """ - emoji_role_id = self.get_rfr_message(guild_id, channel_id, message_id) - if not emoji_role_id: +@assign_session +def get_guild_rfr_roles(guild_id: int) -> List[int]: + """ + Returns all role IDs of roles given by RFR messages in a guild + + :param guild_id: Guild ID to check in. + :return: Role IDs of RFR roles in a specific guild + """ + with session_manager() as session: + rfr_messages = session.execute(select(GuildRFRMessages).filter_by(guild_id=guild_id)).scalars().all() + if not rfr_messages: + return [] + role_ids: List[int] = [] + for rfr_message in rfr_messages: + roles: List[Tuple[int, str, int]] = get_rfr_message_emoji_roles(rfr_message.emoji_role_id) + if not roles: + continue + ids: List[int] = [x[2] for x in roles] + role_ids.extend(ids) + return role_ids + + +@assign_session +def get_rfr_message_emoji_roles(emoji_role_id: int, *, session: sqlalchemy.orm.Session): + """ + Returns all the emoji-role combinations on an rfr message + + :param emoji_role_id: emoji-role combo identifier + :param session: + :return: List of rows in the database if found, otherwise None + """ + rows = session.execute(select(RFRMessageEmojiRoles).filter_by(emoji_role_id=emoji_role_id)).scalars().all() + + return [(row.emoji_role_id, row.emoji_raw, row.role_id) for row in rows] + + +@assign_session +def get_rfr_reaction_role(emoji_role_id: int, emoji_raw: str, role_id: int, *, session: sqlalchemy.orm.Session): + """ + Returns a specific emoji-role combo on an rfr message + + :param emoji_role_id: emoji-role combo identifier + :param emoji_raw: raw string representation of the emoji + :param role_id: role ID of the emoji-role combo + :param session: + :return: Unique row corresponding to a specific emoji-role combo + """ + row = session.execute(select(RFRMessageEmojiRoles).filter_by( + emoji_role_id=emoji_role_id, emoji_raw=emoji_raw, role_id=role_id)).scalar() + if row: + return row.emoji_role_id, row.emoji_raw, row.role_id + else: + return None + + +@assign_session +def get_rfr_reaction_role_by_emoji_str(emoji_role_id: int, emoji_raw: str, session: sqlalchemy.orm.Session): + """ + Gets a role ID from the emoji_role_id and the emoji associated with that role in the emoji-role combo + :param emoji_role_id: emoji-role combo identifier + :param emoji_raw: raw string representation of the emoji + :return: role ID of the emoji-role combo + """ + with session_manager() as session: + row = session.execute(select(RFRMessageEmojiRoles.role_id) + .filter_by(emoji_role_id=emoji_role_id, emoji_raw=emoji_raw)).one_or_none() + if not row: return - else: - self.remove_rfr_message_emoji_roles(emoji_role_id[3]) - - with session_manager() as session: - delete_sql = delete(GuildRFRMessages) \ - .where(and_(and_( - GuildRFRMessages.guild_id == guild_id, - GuildRFRMessages.channel_id == channel_id), - GuildRFRMessages.message_id == message_id)) - session.execute(delete_sql) - session.commit() - - def get_rfr_message(self, guild_id: int, channel_id: int, message_id: int) -> Optional[Tuple[int, int, int, int]]: - """ - Gets the unique rfr message that is specified by the guild ID, channel ID and message ID. - :param guild_id: Guild ID of the rfr message - :param channel_id: Channel ID of the rfr message - :param message_id: Message ID of the rfr message - :return: RFR message info of the specific message if found, otherwise None. - """ - with session_manager() as session: - message = session.execute(select(GuildRFRMessages) - .filter_by(guild_id=guild_id, - channel_id=channel_id, - message_id=message_id)).scalars().one_or_none() - if message: - return message.old_format() - else: - return None - - def get_guild_rfr_messages(self, guild_id: int): - """ - Gets all rfr messages in a given guild, from the guild ID - :param guild_id: ID of the guild - :return: List of rfr messages in the guild. - """ - with session_manager() as session: - messages = session.execute(select(GuildRFRMessages) - .filter_by(guild_id=guild_id)).scalars().all() - return [message.old_format() - for message in messages] - - def get_guild_rfr_roles(self, guild_id: int) -> List[int]: - """ - Returns all role IDs of roles given by RFR messages in a guild - - :param guild_id: Guild ID to check in. - :return: Role IDs of RFR roles in a specific guild - """ - with session_manager() as session: - rfr_messages = session.execute(select(GuildRFRMessages).filter_by(guild_id=guild_id)).scalars().all() - if not rfr_messages: - return [] - role_ids: List[int] = [] - for rfr_message in rfr_messages: - roles: List[Tuple[int, str, int]] = self.get_rfr_message_emoji_roles(rfr_message.emoji_role_id) - if not roles: - continue - ids: List[int] = [x[2] for x in roles] - role_ids.extend(ids) - return role_ids - - def get_rfr_message_emoji_roles(self, emoji_role_id: int): - """ - Returns all the emoji-role combinations on an rfr message - - :param emoji_role_id: emoji-role combo identifier - :return: List of rows in the database if found, otherwise None - """ - with session_manager() as session: - rows = session.execute(select(RFRMessageEmojiRoles).filter_by(emoji_role_id=emoji_role_id)).scalars().all() - - return [(row.emoji_role_id, row.emoji_raw, row.role_id) for row in rows] - - def get_rfr_reaction_role(self, emoji_role_id: int, emoji_raw: str, role_id: int): - """ - Returns a specific emoji-role combo on an rfr message - - :param emoji_role_id: emoji-role combo identifier - :param emoji_raw: raw string representation of the emoji - :param role_id: role ID of the emoji-role combo - :return: Unique row corresponding to a specific emoji-role combo - """ - with session_manager() as session: - row = session.execute(select(RFRMessageEmojiRoles).filter_by( - emoji_role_id=emoji_role_id, emoji_raw=emoji_raw, role_id=role_id)).scalar() - if row: - return row.emoji_role_id, row.emoji_raw, row.role_id - else: - return None - - def get_rfr_reaction_role_by_emoji_str(self, emoji_role_id: int, emoji_raw: str) -> Optional[int]: - """ - Gets a role ID from the emoji_role_id and the emoji associated with that role in the emoji-role combo - :param emoji_role_id: emoji-role combo identifier - :param emoji_raw: raw string representation of the emoji - :return: role ID of the emoji-role combo - """ - with session_manager() as session: - row = session.execute(select(RFRMessageEmojiRoles.role_id) - .filter_by(emoji_role_id=emoji_role_id, emoji_raw=emoji_raw)).one_or_none() - if not row: - return - return row[0] - - def add_guild_rfr_required_role(self, guild_id: int, role_id: int): - """ - Adds a role to the list of roles required to use rfr functionality in a guild. - :param guild_id: guild ID - :param role_id: role ID - :return: - """ - with session_manager() as session: - session.add(GuildRFRRequiredRoles(guild_id=guild_id, role_id=role_id)) - session.commit() - - def remove_guild_rfr_required_role(self, guild_id: int, role_id: int): - """ - Removes a role from the list of roles required to use rfr functionality in a guild - :param guild_id: guild ID - :param role_id: role ID - :return: - """ - with session_manager() as session: - session.execute(delete(GuildRFRRequiredRoles).filter_by(guild_id=guild_id, role_id=role_id)) - session.commit() - - def get_guild_rfr_required_roles(self, guild_id) -> List[int]: - """ - Gets the list of role IDs of roles required to use rfr functionality in a guild - :param guild_id: guild ID - :return: List of role IDs - """ - with session_manager() as session: - rows = session.execute(select(GuildRFRRequiredRoles).filter_by(guild_id=guild_id)).scalars().all() - - role_ids = [x.role_id for x in rows] - if not role_ids: - return [] - return role_ids + return row[0] + + +@assign_session +def add_guild_rfr_required_role(guild_id: int, role_id: int, session: sqlalchemy.orm.Session): + """ + Adds a role to the list of roles required to use rfr functionality in a guild. + :param guild_id: guild ID + :param role_id: role ID + :return: + """ + session.add(GuildRFRRequiredRoles(guild_id=guild_id, role_id=role_id)) + session.commit() + + +@assign_session +def remove_guild_rfr_required_role(guild_id: int, role_id: int, session: sqlalchemy.orm.Session): + """ + Removes a role from the list of roles required to use rfr functionality in a guild + :param guild_id: guild ID + :param role_id: role ID + :return: + """ + session.execute(delete(GuildRFRRequiredRoles).filter_by(guild_id=guild_id, role_id=role_id)) + session.commit() + + +@assign_session +def get_guild_rfr_required_roles(guild_id, session: sqlalchemy.orm.Session) -> List[int]: + """ + Gets the list of role IDs of roles required to use rfr functionality in a guild + :param guild_id: guild ID + :return: List of role IDs + """ + rows = session.execute(select(GuildRFRRequiredRoles).filter_by(guild_id=guild_id)).scalars().all() + + role_ids = [x.role_id for x in rows] + if not role_ids: + return [] + return role_ids diff --git a/koala/cogs/react_for_role/dto.py b/koala/cogs/react_for_role/dto.py new file mode 100644 index 00000000..5b1eeab0 --- /dev/null +++ b/koala/cogs/react_for_role/dto.py @@ -0,0 +1,32 @@ +from dataclasses import dataclass +from typing import List + +import discord + + +@dataclass +class ReactRole: + emoji: str + role_id: int + + def to_tuple(self, guild: discord.Guild): + return self.emoji, guild.get_role(self.role_id) + + +@dataclass +class ReactMessage: + message_id: int + guild_id: int + channel_id: int + title: str + description: str + colour: str + thumbnail: str + inline: bool + roles: List[ReactRole] + + +@dataclass +class RequiredRoles: + guild_id: int + role_ids: List[int] diff --git a/koala/rest/api.py b/koala/rest/api.py index 8f13d598..e718bef9 100644 --- a/koala/rest/api.py +++ b/koala/rest/api.py @@ -7,8 +7,13 @@ # Libs from functools import wraps +from typing import OrderedDict + import aiohttp.web +from aiohttp.abc import Request +from aiohttp.typedefs import Handler +from koala.log import logger # Own modules from koala.models import BaseModel @@ -16,6 +21,7 @@ from http.client import OK + # Variables @@ -23,6 +29,7 @@ class EnhancedJSONEncoder(json.JSONEncoder): """ A custom JSON encoder for datatypes used for this project """ + def default(self, o): if isinstance(o, BaseModel): return o.as_dict() @@ -46,11 +53,11 @@ def build_response(status_code, data): body = None return aiohttp.web.Response(status=status_code, - body=body, - content_type='application/json') + body=body, + content_type='application/json') -def parse_request(*args, **kwargs): +def parse_request(*args, **kwargs) -> Handler: """ A wrapper for API endpoints that provide the required args if raw_response = true, then the default response type is not applied @@ -89,33 +96,41 @@ def parsed_request(func): @wraps(func) async def wrapper(*args, **kwargs): self = args[0] - request = args[1] + request: Request = args[1] + + wanted_args: dict[str, inspect.Parameter] = dict(inspect.signature(func).parameters) + wanted_args.pop("self") - wanted_args = list(inspect.signature(func).parameters.keys()) - wanted_args.remove("self") + required_args: dict[str, inspect.Parameter] = {a: wanted_args.get(a) for a in wanted_args.keys() if + wanted_args.get(a).default == inspect.Parameter.empty} available_args = {} - if (request.method == "POST" or request.method == "PUT") and request.has_body: + if (request.method in request.POST_METHODS) and request.can_read_body: body = await request.json() - for arg in wanted_args: + for arg in wanted_args.keys(): if arg in body: available_args[arg] = body[arg] else: - for arg in wanted_args: + for arg in wanted_args.keys(): if arg in request.query: available_args[arg] = request.query[arg] - unsatisfied_args = set(wanted_args) - set(available_args.keys()) + unsatisfied_args = set(required_args.keys()) - set(available_args.keys()) if unsatisfied_args: # Expected match info that doesn't exist raise aiohttp.web.HTTPBadRequest(reason="Unsatisfied Arguments: %s" % unsatisfied_args) - result = await func(self, **{arg_name: available_args[arg_name] for arg_name in wanted_args}) + try: + result = await func(self, **{arg_name: available_args[arg_name] for arg_name in available_args.keys()}) + except Exception as e: + logger.error("API Failed", exc_info=e) + raise e if raw_response: return result else: return build_response(OK, result) return wrapper + return parsed_request(func) if func else parsed_request diff --git a/requirements.txt b/requirements.txt index d5f25076..178071eb 100644 --- a/requirements.txt +++ b/requirements.txt @@ -7,7 +7,7 @@ certifi==2022.12.7 chardet==5.1.0 colorama==0.4.6 discord.py==2.2.2 -dpytest==0.6.3 +dpytest==0.6.4 emoji==1.7.0 idna==3.4 iniconfig==2.0.0 diff --git a/tests/cogs/base/test_api.py b/tests/cogs/base/test_api.py index d1d71dcd..c0a67842 100644 --- a/tests/cogs/base/test_api.py +++ b/tests/cogs/base/test_api.py @@ -53,67 +53,67 @@ async def test_get_activities_missing_param(api_client): ''' async def test_put_schedule_activity(api_client): - resp = await api_client.put('/scheduled-activity', json=( + resp = await api_client.put('/scheduled-activity', json= { 'activity_type': 'playing', 'message': 'test', 'url': 'test.com', 'start_time': '2025-01-01 00:00:00', 'end_time': '2026-01-01 00:00:00' - })) + }) assert resp.status == CREATED text = await resp.text() assert text == '{"message": "Activity scheduled"}' async def test_put_schedule_activity_missing_param(api_client): - resp = await api_client.put('/scheduled-activity', json=( + resp = await api_client.put('/scheduled-activity', json= { 'activity_type': 'playing', 'message': 'test', 'url': 'test.com', 'start_time': '2025-01-01 00:00:00' - })) + }) assert resp.status == BAD_REQUEST text = await resp.text() assert text == "400: Unsatisfied Arguments: {'end_time'}" async def test_put_schedule_activity_bad_activity(api_client): - resp = await api_client.put('/scheduled-activity', json=( + resp = await api_client.put('/scheduled-activity', json= { 'activity_type': 'invalidActivity', 'message': 'test', 'url': 'test.com', 'start_time': '2025-01-01 00:00:00', 'end_time': '2026-01-01 00:00:00' - })) + }) assert resp.status == UNPROCESSABLE_ENTITY assert await resp.text() == '422: Error scheduling activity: Invalid activity type' async def test_put_schedule_activity_bad_start_time(api_client): - resp = await api_client.put('/scheduled-activity', json=( + resp = await api_client.put('/scheduled-activity', json= { 'activity_type': 'playing', 'message': 'test', 'url': 'test.com', 'start_time': 'invalid_time', 'end_time': '2026-01-01 00:00:00' - })) + }) assert resp.status == UNPROCESSABLE_ENTITY assert await resp.text() == '422: Error scheduling activity: Bad start / end time' async def test_put_schedule_activity_bad_end_time(api_client): - resp = await api_client.put('/scheduled-activity', json=( + resp = await api_client.put('/scheduled-activity', json= { 'activity_type': 'invalidActivity', 'message': 'test', 'url': 'test.com', 'start_time': '2026-01-01 00:00:00', 'end_time': 'invalidTime' - })) + }) assert resp.status == UNPROCESSABLE_ENTITY assert await resp.text() == '422: Error scheduling activity: Bad start / end time' @@ -125,12 +125,12 @@ async def test_put_schedule_activity_bad_end_time(api_client): async def test_put_set_activity(api_client): - resp = await api_client.put('/activity', json=( + resp = await api_client.put('/activity', json= { 'activity_type': 'playing', 'name': 'test', 'url': 'test.com' - })) + }) assert resp.status == CREATED text = await resp.text() assert text == '{"message": "Activity set"}' @@ -138,22 +138,22 @@ async def test_put_set_activity(api_client): async def test_put_set_activity_bad_req(api_client): - resp = await api_client.put('/activity', json=( + resp = await api_client.put('/activity', json= { 'activity_type': 'invalidActivity', 'name': 'test', 'url': 'test.com' - })) + }) assert resp.status == UNPROCESSABLE_ENTITY assert await resp.text() == '422: Error setting activity: Invalid activity type' async def test_put_set_activity_missing_param(api_client): - resp = await api_client.put('/activity', json=( + resp = await api_client.put('/activity', json= { 'activity_type': 'invalidActivity', 'url': 'test.com' - })) + }) assert resp.status == BAD_REQUEST assert await resp.text() == "400: Unsatisfied Arguments: {'name'}" @@ -203,54 +203,51 @@ async def test_get_support_link(api_client): ''' async def test_post_load_cog(api_client): - resp = await api_client.post('/load-cog', json=( + resp = await api_client.post('/load-cog', json= { 'extension': 'announce', - # 'package': koalabot.COGS_PACKAGE - })) + 'package': koalabot.COGS_PACKAGE + }) assert resp.status == OK text = await resp.text() assert text == '{"message": "Cog loaded"}' async def test_post_load_base_cog(api_client): - resp = await api_client.post('/load-cog', json=( + resp = await api_client.post('/load-cog', json= { 'extension': 'base', - # 'package': koalabot.COGS_PACKAGE - })) + 'package': koalabot.COGS_PACKAGE + }) assert resp.status == OK text = await resp.text() assert text == '{"message": "Cog loaded"}' async def test_post_load_cog_bad_req(api_client): - resp = await api_client.post('/load-cog', json=( + resp = await api_client.post('/load-cog', json= { 'extension': 'invalidCog', - # 'package': koalabot.COGS_PACKAGE - })) + 'package': koalabot.COGS_PACKAGE + }) assert resp.status == UNPROCESSABLE_ENTITY assert await resp.text() == '422: Error loading cog: Invalid extension' -# async def test_post_load_cog_missing_param(api_client): -# resp = await api_client.post('/load-cog', json=( -# { -# 'extension': 'invalidCog' -# })) -# assert resp.status == BAD_REQUEST -# assert await resp.text() == "400: Unsatisfied Arguments: {'package'}" +async def test_post_load_cog_missing_param(api_client): + resp = await api_client.post('/load-cog', json={}) + assert resp.status == BAD_REQUEST + assert await resp.text() == "400: Unsatisfied Arguments: {'extension'}" async def test_post_load_cog_already_loaded(api_client): - await api_client.post('/load-cog', json=( + await api_client.post('/load-cog', json= { 'extension': 'announce', - # 'package': koalabot.COGS_PACKAGE - })) + 'package': koalabot.COGS_PACKAGE + }) - resp = await api_client.post('/load-cog', json=( + resp = await api_client.post('/load-cog', json= { 'extension': 'announce', - # 'package': koalabot.COGS_PACKAGE - })) + 'package': koalabot.COGS_PACKAGE + }) assert resp.status == UNPROCESSABLE_ENTITY assert await resp.text() == '422: Error loading cog: Already loaded' @@ -261,44 +258,41 @@ async def test_post_load_cog_already_loaded(api_client): ''' async def test_post_unload_cog(api_client): - await api_client.post('/load-cog', json=( + await api_client.post('/load-cog', json= { 'extension': 'announce', - # 'package': koalabot.COGS_PACKAGE - })) + 'package': koalabot.COGS_PACKAGE + }) - resp = await api_client.post('/unload-cog', json=( + resp = await api_client.post('/unload-cog', json= { 'extension': 'announce', - # 'package': koalabot.COGS_PACKAGE - })) + 'package': koalabot.COGS_PACKAGE + }) assert resp.status == OK text = await resp.text() assert text == '{"message": "Cog unloaded"}' async def test_post_unload_cog_not_loaded(api_client): - resp = await api_client.post('/unload-cog', json=( + resp = await api_client.post('/unload-cog', json= { 'extension': 'announce', - # 'package': koalabot.COGS_PACKAGE - })) + 'package': koalabot.COGS_PACKAGE + }) assert resp.status == UNPROCESSABLE_ENTITY assert await resp.text() == '422: Error unloading cog: Extension not loaded' -# async def test_post_unload_cog_missing_param(api_client): -# resp = await api_client.post('/unload-cog', json=( -# { -# 'extension': 'invalidCog' -# })) -# assert resp.status == BAD_REQUEST -# assert await resp.text() == "400: Unsatisfied Arguments: {'package'}" +async def test_post_unload_cog_missing_param(api_client): + resp = await api_client.post('/unload-cog', json={}) + assert resp.status == BAD_REQUEST + assert await resp.text() == "400: Unsatisfied Arguments: {'extension'}" async def test_post_unload_base_cog(api_client): - resp = await api_client.post('/unload-cog', json=( + resp = await api_client.post('/unload-cog', json= { 'extension': 'BaseCog', - # 'package': koalabot.COGS_PACKAGE - })) + 'package': koalabot.COGS_PACKAGE + }) assert resp.status == UNPROCESSABLE_ENTITY text = await resp.text() assert text == "422: Error unloading cog: Sorry, you can't unload the base cog" @@ -313,10 +307,10 @@ async def test_post_unload_base_cog(api_client): async def test_post_enable_extension(api_client, bot): await koalabot.load_all_cogs(bot) guild: discord.Guild = dpytest.get_config().guilds[0] - resp = await api_client.post('/enable-extension', json=({ + resp = await api_client.post('/enable-extension', json={ 'guild_id': guild.id, 'koala_ext': 'Announce' - })) + }) assert resp.status == OK text = await resp.text() @@ -325,21 +319,21 @@ async def test_post_enable_extension(api_client, bot): async def test_post_enable_extension_bad_req(api_client): guild: discord.Guild = dpytest.get_config().guilds[0] - resp = await api_client.post('/enable-extension', json=( + resp = await api_client.post('/enable-extension', json= { 'guild_id': guild.id, 'koala_ext': 'Invalid Extension' - })) + }) assert resp.status == UNPROCESSABLE_ENTITY text = await resp.text() assert text == "422: Error enabling extension: Invalid extension" async def test_post_enable_extension_missing_param(api_client): guild: discord.Guild = dpytest.get_config().guilds[0] - resp = await api_client.post('/enable-extension', json=( + resp = await api_client.post('/enable-extension', json= { 'guild_id': guild.id - })) + }) assert resp.status == BAD_REQUEST text = await resp.text() assert text == "400: Unsatisfied Arguments: {'koala_ext'}" @@ -354,35 +348,35 @@ async def test_post_enable_extension_missing_param(api_client): async def test_post_disable_extension(api_client, bot): await koalabot.load_all_cogs(bot) guild: discord.Guild = dpytest.get_config().guilds[0] - setup = await api_client.post('/enable-extension', json=({ + setup = await api_client.post('/enable-extension', json={ 'guild_id': guild.id, 'koala_ext': 'Announce' - })) + }) assert setup.status == OK - resp = await api_client.post('/disable-extension', json=({ + resp = await api_client.post('/disable-extension', json={ 'guild_id': guild.id, 'koala_ext': 'Announce' - })) + }) assert resp.status == OK text = await resp.text() assert text == '{"message": "Extension disabled"}' async def test_post_disable_extension_not_enabled(api_client): guild: discord.Guild = dpytest.get_config().guilds[0] - resp = await api_client.post('/disable-extension', json=({ + resp = await api_client.post('/disable-extension', json={ 'guild_id': guild.id, 'koala_ext': 'Announce' - })) + }) assert resp.status == UNPROCESSABLE_ENTITY text = await resp.text() assert text == "422: Error disabling extension: Extension not enabled" async def test_post_disable_extension_missing_param(api_client): guild: discord.Guild = dpytest.get_config().guilds[0] - resp = await api_client.post('/disable-extension', json=({ + resp = await api_client.post('/disable-extension', json={ 'guild_id': guild.id - })) + }) assert resp.status == BAD_REQUEST text = await resp.text() assert text == "400: Unsatisfied Arguments: {'koala_ext'}" @@ -390,11 +384,11 @@ async def test_post_disable_extension_missing_param(api_client): async def test_post_disable_extension_bad_req(api_client): guild: discord.Guild = dpytest.get_config().guilds[0] - resp = await api_client.post('/disable-extension', json=( + resp = await api_client.post('/disable-extension', json= { 'guild_id': guild.id, 'koala_ext': 'Invalid Extension' - })) + }) assert resp.status == UNPROCESSABLE_ENTITY text = await resp.text() assert text == "422: Error disabling extension: Extension not enabled" diff --git a/tests/cogs/react_for_role/test_api.py b/tests/cogs/react_for_role/test_api.py new file mode 100644 index 00000000..36dee8d9 --- /dev/null +++ b/tests/cogs/react_for_role/test_api.py @@ -0,0 +1,234 @@ +from http.client import BAD_REQUEST, CREATED, OK, UNPROCESSABLE_ENTITY + +from mock import mock +from koala.db import get_all_available_guild_extensions +from koala.rest.api import parse_request + +import koalabot +from koala.cogs.react_for_role.api import RfrEndpoint, MESSAGE, REQUIRED_ROLES + +# Libs +import discord +from aiohttp import web +import pytest +import discord.ext.test as dpytest + + +@pytest.fixture +def api_client(bot: discord.ext.commands.Bot, aiohttp_client, loop): + app = web.Application() + endpoint = RfrEndpoint(bot) + app = endpoint.register(app) + return loop.run_until_complete(aiohttp_client(app)) + + +async def test_message_post_partial(api_client): + resp = await api_client.post('/{}'.format(MESSAGE), json={ + "guild_id": dpytest.get_config().guilds[0].id, + "channel_id": dpytest.get_config().guilds[0].channels[0].id, + "title": "API test", + "description": "desc", + "colour": "#0000ff" + }) + assert resp.status == OK + resp_json: dict = await resp.json() + assert "message_id" in resp_json.keys() + + +async def test_message_post_full(api_client): + resp = await api_client.post('/{}'.format(MESSAGE), json={ + "guild_id": dpytest.get_config().guilds[0].id, + "channel_id": dpytest.get_config().guilds[0].channels[0].id, + "title": "API test", + "description": "desc", + "colour": "#0000ff", + "thumbnail": "https://koalabot.uk/static/media/KoalaBotLogo-min.78f6a0d317dfdfa7391d.png", + "inline": "true", + "roles": [{ + "role_id": dpytest.get_config().guilds[0].roles[0].id, + "emoji": "<:discordmod:1030226250884722809>" + }] + }) + assert resp.status == OK + resp_json: dict = await resp.json() + assert "message_id" in resp_json.keys() + + +async def test_message_get(api_client): + resp1 = await api_client.post('/{}'.format(MESSAGE), json={ + "guild_id": dpytest.get_config().guilds[0].id, + "channel_id": dpytest.get_config().guilds[0].channels[0].id, + "title": "API test", + "description": "desc", + "colour": "#0000ff" + }) + message_id = (await resp1.json())["message_id"] + + resp = await api_client.get('/{}?message_id={}&guild_id={}&channel_id={}' + .format(MESSAGE, + message_id, + dpytest.get_config().guilds[0].id, + dpytest.get_config().guilds[0].channels[0].id)) + + assert resp.status == OK + resp_json: dict = await resp.json() + assert resp_json.get("colour") == "#0000ff" + + +async def test_message_put(api_client): + resp1 = await api_client.post('/{}'.format(MESSAGE), json={ + "guild_id": dpytest.get_config().guilds[0].id, + "channel_id": dpytest.get_config().guilds[0].channels[0].id, + "title": "API test", + "description": "desc", + "colour": "#0000ff" + }) + post_response = await resp1.json() + post_response["colour"] = "#ffffff" + post_response["title"] = "test2" + post_response["description"] = "desc2" + assert post_response["thumbnail"] == "https://cdn.discordapp.com/attachments/737280260541907015/752024535985029240/discord1.png" + post_response["thumbnail"] = "https://koalabot.uk/static/media/KoalaBotLogo-min.78f6a0d317dfdfa7391d.png" + assert post_response["inline"] is False + post_response["inline"] = True + assert post_response["roles"] == [] + post_response["roles"] = [{ + "role_id": dpytest.get_config().guilds[0].roles[0].id, + "emoji": "<:discordmod:1030226250884722809>" + }] + resp = await api_client.put('/{}'.format(MESSAGE), json=post_response) + + assert resp.status == OK + resp_json: dict = await resp.json() + assert resp_json.get("colour") == "#ffffff" + assert resp_json.get("title") == "test2" + assert resp_json.get("description") == "desc2" + assert resp_json["thumbnail"] == "https://koalabot.uk/static/media/KoalaBotLogo-min.78f6a0d317dfdfa7391d.png" + assert resp_json["inline"] is True + assert resp_json["roles"] == [{ + "role_id": dpytest.get_config().guilds[0].roles[0].id, + "emoji": "<:discordmod:1030226250884722809>" + }] + + +async def test_message_patch_partial(api_client): + resp1 = await api_client.post('/{}'.format(MESSAGE), json={ + "guild_id": dpytest.get_config().guilds[0].id, + "channel_id": dpytest.get_config().guilds[0].channels[0].id, + "title": "API test", + "description": "desc", + "colour": "#0000ff" + }) + post_response = await resp1.json() + + patch_body = { + "message_id": post_response["message_id"], + "guild_id": dpytest.get_config().guilds[0].id, + "channel_id": dpytest.get_config().guilds[0].channels[0].id, + "description": "desc2" + } + resp = await api_client.patch('/{}'.format(MESSAGE), json=patch_body) + + assert resp.status == OK + message = await dpytest.get_config().guilds[0].channels[0].fetch_message(post_response["message_id"]) + assert message.embeds[0].description == "desc2" + resp_json: dict = await resp.json() + assert resp_json.get("colour") == "#0000ff" + assert resp_json.get("description") == "desc2" + + +async def test_message_patch_full(api_client): + resp1 = await api_client.post('/{}'.format(MESSAGE), json={ + "guild_id": dpytest.get_config().guilds[0].id, + "channel_id": dpytest.get_config().guilds[0].channels[0].id, + "title": "API test", + "description": "desc", + "colour": "#0000ff" + }) + post_response = await resp1.json() + + patch_body = { + "message_id": post_response["message_id"], + "guild_id": dpytest.get_config().guilds[0].id, + "channel_id": dpytest.get_config().guilds[0].channels[0].id, + "title": "test2", + "description": "desc2", + "colour": "#000fff", + "thumbnail": "https://koalabot.uk/static/media/KoalaBotLogo-min.78f6a0d317dfdfa7391d.png", + "inline": "true", + "roles": [{ + "role_id": dpytest.get_config().guilds[0].roles[0].id, + "emoji": "<:discordmod:1030226250884722809>" + }] + } + resp = await api_client.patch('/{}'.format(MESSAGE), json=patch_body) + + assert resp.status == OK + message = await dpytest.get_config().guilds[0].channels[0].fetch_message(post_response["message_id"]) + assert message.embeds[0].description == "desc2" + resp_json: dict = await resp.json() + assert resp_json.get("colour") == "#000fff" + assert resp_json.get("title") == "test2" + assert resp_json.get("description") == "desc2" + assert resp_json["thumbnail"] == "https://koalabot.uk/static/media/KoalaBotLogo-min.78f6a0d317dfdfa7391d.png" + assert resp_json["inline"] is True + assert resp_json["roles"] == [{ + "role_id": dpytest.get_config().guilds[0].roles[0].id, + "emoji": "<:discordmod:1030226250884722809>" + }] + + +async def test_message_delete(api_client): + resp1 = await api_client.post('/{}'.format(MESSAGE), json={ + "guild_id": dpytest.get_config().guilds[0].id, + "channel_id": dpytest.get_config().guilds[0].channels[0].id, + "title": "API test", + "description": "desc", + "colour": "#0000ff" + }) + post_response = await resp1.json() + + message = await dpytest.get_config().guilds[0].channels[0].fetch_message(post_response["message_id"]) + assert message is not None + assert message.embeds[0].description == "desc" + + delete_body = { + "message_id": post_response["message_id"], + "guild_id": dpytest.get_config().guilds[0].id, + "channel_id": dpytest.get_config().guilds[0].channels[0].id + } + resp = await api_client.delete('/{}'.format(MESSAGE), json=delete_body) + + assert resp.status == OK + with pytest.raises(discord.NotFound): + await dpytest.get_config().guilds[0].channels[0].fetch_message(post_response["message_id"]) + resp_json: dict = await resp.json() + assert resp_json.get("status") == "DELETED" + assert resp_json.get("message_id") == post_response["message_id"] + + +# /REQUIRED_ROLES + +async def test_required_roles_put(api_client): + resp = await api_client.put('/{}'.format(REQUIRED_ROLES), json={ + "guild_id": dpytest.get_config().guilds[0].id, + "role_ids": [dpytest.get_config().guilds[0].roles[0].id] + }) + assert resp.status == OK + resp_json: dict = await resp.json() + assert resp_json.get("role_ids") == [dpytest.get_config().guilds[0].roles[0].id] + assert resp_json.get("guild_id") == dpytest.get_config().guilds[0].id + + +async def test_required_roles_get(api_client): + await api_client.put('/{}'.format(REQUIRED_ROLES), json={ + "guild_id": dpytest.get_config().guilds[0].id, + "role_ids": [dpytest.get_config().guilds[0].roles[0].id] + }) + + resp = await api_client.get('/{}?guild_id={}'.format(REQUIRED_ROLES, dpytest.get_config().guilds[0].id)) + + assert resp.status == OK + resp_json: dict = await resp.json() + assert resp_json.get("role_ids") == [dpytest.get_config().guilds[0].roles[0].id] + assert resp_json.get("guild_id") == dpytest.get_config().guilds[0].id diff --git a/tests/cogs/react_for_role/test_cog.py b/tests/cogs/react_for_role/test_cog.py index 0c671ddc..4a72d127 100644 --- a/tests/cogs/react_for_role/test_cog.py +++ b/tests/cogs/react_for_role/test_cog.py @@ -5,8 +5,6 @@ Commented using reStructuredText (reST) """ -# Futures - # Built-in/Generic Imports import random @@ -21,13 +19,17 @@ from discord.http import MultipartParameters # Own modules +from koala.cogs.react_for_role import core import koalabot from koala.cogs import ReactForRole from koala.colours import KOALA_GREEN from koala.db import session_manager from tests.tests_utils import utils as testutils -from .utils import DBManager, independent_get_guild_rfr_message, independent_get_guild_rfr_required_role +from koala.cogs.react_for_role.db import * +from .utils import independent_get_guild_rfr_message, independent_get_guild_rfr_required_role from tests.log import logger +from koala.cogs import ReactForRole + # Constants @@ -58,7 +60,7 @@ async def test_get_rfr_message_from_prompts(bot, utils_cog, rfr_cog): await rfr_cog.get_rfr_message_from_prompts(ctx) assert str( exc.value) == "Message ID given is not that of a react for role message." - DBManager.add_rfr_message(msg.guild.id, channel_id, msg_id) + add_rfr_message(msg.guild.id, channel_id, msg_id) with mock.patch('koala.cogs.ReactForRole.prompt_for_input', side_effect=[str(channel_id), str(msg_id)]) as mock_input: with mock.patch('discord.abc.Messageable.fetch_message', mock.AsyncMock(return_value=msg)): @@ -67,7 +69,6 @@ async def test_get_rfr_message_from_prompts(bot, utils_cog, rfr_cog): assert rfr_msg_channel.id == channel_id -# TODO Actually implement the test. @pytest.mark.parametrize("num_rows", [0, 1, 2, 20, 100, 250]) @pytest.mark.asyncio async def test_parse_emoji_and_role_input_str(num_rows, utils_cog, rfr_cog): @@ -201,10 +202,12 @@ async def test_overwrite_channel_add_reaction_perms(rfr_cog: ReactForRole): config: dpytest.RunnerConfig = dpytest.get_config() guild: discord.Guild = config.guilds[0] channel: discord.TextChannel = guild.text_channels[0] + bot: discord.Client = config.client with mock.patch('discord.ext.test.backend.FakeHttp.edit_channel_permissions') as mock_edit_channel_perms: for i in range(15): await guild.create_role(name=f"TestRole{i}", permissions=discord.Permissions.all()) role: discord.Role = discord.utils.get(guild.roles, id=guild.id) + # await core.setup_rfr_reaction_permissions(guild, channel, bot) await rfr_cog.overwrite_channel_add_reaction_perms(guild, channel) calls = [mock.call(channel.id, role.id, '0', '64', discord.abc._Overwrites.ROLE, reason=None), mock.call(channel.id, config.client.user.id, '64', '0', discord.abc._Overwrites.MEMBER, @@ -227,21 +230,20 @@ async def test_is_user_alive(utils_cog, rfr_cog): @pytest.mark.asyncio -async def test_get_embed_from_message(rfr_cog): +async def test_get_embed_from_message(rfr_cog, bot: commands.Bot): config: dpytest.RunnerConfig = dpytest.get_config() author: discord.Member = config.members[0] guild: discord.Guild = config.guilds[0] channel: discord.TextChannel = guild.text_channels[0] - test_embed_dict: dict = {'title': 'title', 'description': 'descr', 'type': 'rich', 'url': 'https://www.google.com'} - bot: discord.Client = config.client - await bot.http.send_message(channel.id, params=MultipartParameters({"embeds": [test_embed_dict]}, None, None)) + embed = discord.Embed(title="title", description="descr", type="rich", url="https://www.google.com") + await channel.send(embed=embed) sent_msg: discord.Message = await dpytest.sent_queue.get() msg_mock: discord.Message = dpytest.back.make_message('a', author, channel) - result = rfr_cog.get_embed_from_message(None) + result = core.get_embed_from_message(None) assert result is None - result = rfr_cog.get_embed_from_message(msg_mock) + result = core.get_embed_from_message(msg_mock) assert result is None - result = rfr_cog.get_embed_from_message(sent_msg) + result = core.get_embed_from_message(sent_msg) assert dpytest.embed_eq(result, sent_msg.embeds[0]) @@ -259,12 +261,12 @@ async def test_get_number_of_embed_fields(rfr_cog): for i in range(20): test_embed.add_field(name=f'field{i}', value=f'num{i}') num_fields += 1 - assert rfr_cog.get_number_of_embed_fields(test_embed) == num_fields + assert core.get_number_of_embed_fields(embed=test_embed) == num_fields @pytest.mark.skip('dpytest currently has non-implemented functionality for construction of guild custom emojis') @pytest.mark.asyncio -async def test_get_first_emoji_from_str(utils_cog, rfr_cog): +async def test_get_first_emoji_from_str(bot, utils_cog, rfr_cog): await dpytest.message(koalabot.COMMAND_PREFIX + "store_ctx") ctx: commands.Context = utils_cog.get_last_ctx() config: dpytest.RunnerConfig = dpytest.get_config() @@ -281,7 +283,7 @@ async def test_get_first_emoji_from_str(utils_cog, rfr_cog): author: discord.Member = config.members[0] channel: discord.TextChannel = guild.text_channels[0] msg: discord.Message = dpytest.back.make_message(str(guild_emoji), author, channel) - result = await rfr_cog.get_first_emoji_from_str(ctx, msg.content) + result = await core.get_first_emoji_from_str(bot, guild, msg.content) logger.debug(result) assert isinstance(result, discord.Emoji), msg.content assert guild_emoji == result @@ -342,7 +344,7 @@ async def test_rfr_delete_message(): channel: discord.TextChannel = guild.text_channels[0] message: discord.Message = await dpytest.message("rfr") msg_id = message.id - DBManager.add_rfr_message(guild.id, channel.id, msg_id) + add_rfr_message(guild.id, channel.id, msg_id) await dpytest.empty_queue() with mock.patch('koala.cogs.ReactForRole.get_rfr_message_from_prompts', mock.AsyncMock(return_value=(message, channel))): @@ -367,13 +369,13 @@ async def test_rfr_edit_description(): client: discord.Client = config.client message: discord.Message = await dpytest.message("rfr") msg_id = message.id - DBManager.add_rfr_message(guild.id, channel.id, msg_id) + add_rfr_message(guild.id, channel.id, msg_id) assert embed.description == 'description' with mock.patch('koala.cogs.ReactForRole.get_rfr_message_from_prompts', mock.AsyncMock(return_value=(message, channel))): with mock.patch('koala.cogs.ReactForRole.prompt_for_input', mock.AsyncMock(side_effect=["new description", "Y"])): - with mock.patch('koala.cogs.ReactForRole.get_embed_from_message', return_value=embed): + with mock.patch('koala.cogs.react_for_role.core.get_embed_from_message', return_value=embed): await dpytest.message(koalabot.COMMAND_PREFIX + "rfr edit description") assert embed.description == 'new description' assert dpytest.verify().message() @@ -390,13 +392,13 @@ async def test_rfr_edit_title(): client: discord.Client = config.client message: discord.Message = await dpytest.message("rfr") msg_id = message.id - DBManager.add_rfr_message(guild.id, channel.id, msg_id) + add_rfr_message(guild.id, channel.id, msg_id) assert embed.title == 'title' with mock.patch('koala.cogs.ReactForRole.get_rfr_message_from_prompts', mock.AsyncMock(return_value=(message, channel))): with mock.patch('koala.cogs.ReactForRole.prompt_for_input', mock.AsyncMock(side_effect=["new title", "Y"])): - with mock.patch('koala.cogs.ReactForRole.get_embed_from_message', return_value=embed): + with mock.patch('koala.cogs.react_for_role.core.get_embed_from_message', return_value=embed): await dpytest.message(koalabot.COMMAND_PREFIX + "rfr edit title") assert embed.title == 'new title' assert dpytest.verify().message() @@ -422,12 +424,12 @@ async def test_rfr_edit_thumbnail_attach(): content_type="image/jpeg")) msg_id = message.id bad_attach = "something that's not an attachment" - DBManager.add_rfr_message(guild.id, channel.id, msg_id) + add_rfr_message(guild.id, channel.id, msg_id) assert embed.thumbnail.url == "https://media.discordapp.net/attachments/611574654502699010/756152703801098280/IMG_20200917_150032.jpg" with mock.patch('koala.cogs.ReactForRole.get_rfr_message_from_prompts', mock.AsyncMock(return_value=(message, channel))): - with mock.patch('koala.cogs.ReactForRole.get_embed_from_message', return_value=embed): + with mock.patch('koala.cogs.react_for_role.core.get_embed_from_message', return_value=embed): with mock.patch('koala.cogs.ReactForRole.prompt_for_input', return_value=attach): await dpytest.message("k!rfr edit image") assert embed.thumbnail.url == "https://media.discordapp.net/attachments/some_number/random_number/test.jpg" @@ -447,12 +449,12 @@ async def test_rfr_edit_thumbnail_bad_attach(attach): url="https://media.discordapp.net/attachments/611574654502699010/756152703801098280/IMG_20200917_150032.jpg") message: discord.Message = await dpytest.message("rfr") msg_id = message.id - DBManager.add_rfr_message(guild.id, channel.id, msg_id) + add_rfr_message(guild.id, channel.id, msg_id) assert embed.thumbnail.url == "https://media.discordapp.net/attachments/611574654502699010/756152703801098280/IMG_20200917_150032.jpg" with mock.patch('koala.cogs.ReactForRole.get_rfr_message_from_prompts', mock.AsyncMock(return_value=(message, channel))): - with mock.patch('koala.cogs.ReactForRole.get_embed_from_message', return_value=embed): + with mock.patch('koala.cogs.react_for_role.core.get_embed_from_message', return_value=embed): with mock.patch('koala.cogs.ReactForRole.prompt_for_input', return_value=attach): with pytest.raises((aiohttp.ClientError, aiohttp.InvalidURL, commands.BadArgument, commands.CommandInvokeError)) as exc: @@ -476,12 +478,12 @@ async def test_rfr_edit_thumbnail_links(image_url): url="https://media.discordapp.net/attachments/611574654502699010/756152703801098280/IMG_20200917_150032.jpg") message: discord.Message = await dpytest.message("rfr") msg_id = message.id - DBManager.add_rfr_message(guild.id, channel.id, msg_id) + add_rfr_message(guild.id, channel.id, msg_id) assert embed.thumbnail.url == "https://media.discordapp.net/attachments/611574654502699010/756152703801098280/IMG_20200917_150032.jpg" with mock.patch('koala.cogs.ReactForRole.get_rfr_message_from_prompts', mock.AsyncMock(return_value=(message, channel))): - with mock.patch('koala.cogs.ReactForRole.get_embed_from_message', return_value=embed): + with mock.patch('koala.cogs.react_for_role.core.get_embed_from_message', return_value=embed): with mock.patch('koala.cogs.ReactForRole.prompt_for_input', return_value=image_url): assert embed.thumbnail.url == "https://media.discordapp.net/attachments/611574654502699010/756152703801098280/IMG_20200917_150032.jpg" await dpytest.message("k!rfr edit image") @@ -503,14 +505,14 @@ async def test_rfr_edit_inline_all(arg): message2: discord.Message = await dpytest.message("rfr") msg1_id = message1.id msg2_id = message2.id - DBManager.add_rfr_message(guild.id, channel.id, msg1_id) - DBManager.add_rfr_message(guild.id, channel.id, msg2_id) + add_rfr_message(guild.id, channel.id, msg1_id) + add_rfr_message(guild.id, channel.id, msg2_id) await dpytest.sent_queue.empty() calls = [mock.call(0, name="field1", value="value1", inline=(arg == "Y")), mock.call(0, name="field2", value="value2", inline=(arg == "Y"))] with mock.patch("koala.cogs.ReactForRole.prompt_for_input", side_effects=["all", arg]): with mock.patch("discord.abc.Messageable.fetch_message", side_effects=[message1, message2]): - with mock.patch("koala.cogs.ReactForRole.get_embed_from_message", side_effects=[embed1, embed2]): + with mock.patch("koala.cogs.react_for_role.core.get_embed_from_message", side_effects=[embed1, embed2]): with mock.patch('discord.Embed.set_field_at') as mock_call: await dpytest.message("k!rfr edit inline") assert dpytest.verify().message() @@ -535,7 +537,7 @@ async def test_rfr_add_roles_to_msg(): author: discord.Member = config.members[0] message: discord.Message = await dpytest.message("rfr") msg_id: int = message.id - DBManager.add_rfr_message(guild.id, channel.id, msg_id) + add_rfr_message(guild.id, channel.id, msg_id) input_em_ro_content = "" em_list = [] ro_list = [] @@ -549,7 +551,7 @@ async def test_rfr_add_roles_to_msg(): with mock.patch('koala.cogs.ReactForRole.get_rfr_message_from_prompts', mock.AsyncMock(return_value=(message, channel))): - with mock.patch('koala.cogs.ReactForRole.get_embed_from_message', return_value=embed): + with mock.patch('koala.cogs.react_for_role.core.get_embed_from_message', return_value=embed): with mock.patch('discord.client.Client.wait_for', mock.AsyncMock(return_value=input_em_ro_msg)): with mock.patch('discord.Embed.add_field') as add_field: @@ -569,7 +571,7 @@ async def test_rfr_remove_roles_from_msg(): author: discord.Member = config.members[0] message: discord.Message = await dpytest.message("rfr") msg_id: int = message.id - DBManager.add_rfr_message(guild.id, channel.id, msg_id) + add_rfr_message(guild.id, channel.id, msg_id) input_em_ro_content = "" em_ro_list = [] for i in range(5): @@ -579,17 +581,17 @@ async def test_rfr_remove_roles_from_msg(): input_em_ro_content += f"{x}\n\r" em_ro_list.append(x) embed.add_field(name=str(em), value=ro.mention, inline=False) - DBManager.add_rfr_message_emoji_role(1, str(em), ro.id) + add_rfr_message_emoji_role(1, str(em), ro.id) input_em_ro_msg: discord.Message = dpytest.back.make_message(input_em_ro_content, author, channel) with mock.patch('koala.cogs.ReactForRole.get_rfr_message_from_prompts', mock.AsyncMock(return_value=(message, channel))): - with mock.patch('koala.cogs.ReactForRole.get_embed_from_message', return_value=embed): + with mock.patch('koala.cogs.react_for_role.core.get_embed_from_message', return_value=embed): with mock.patch('discord.client.Client.wait_for', mock.AsyncMock(return_value=input_em_ro_msg)): with mock.patch('discord.Embed.add_field') as add_field: with mock.patch( - 'koala.cogs.react_for_role.db.ReactForRoleDBManager.remove_rfr_message_emoji_role') as remove_emoji_role: + 'koala.cogs.react_for_role.db.remove_rfr_message_emoji_role') as remove_emoji_role: add_field.reset_mock() await dpytest.message(koalabot.COMMAND_PREFIX + "rfr removeRoles") add_field.assert_not_called() @@ -613,7 +615,7 @@ async def test_can_have_rfr_role(num_roles, num_required, rfr_cog): r_list.append(role) required = random.sample(list(r_list), num_required) for r in required: - DBManager.add_guild_rfr_required_role(guild.id, r.id) + add_guild_rfr_required_role(guild.id, r.id) assert independent_get_guild_rfr_required_role(session, guild.id, r.id) is not None for i in range(num_roles): mem_roles = [] @@ -628,3 +630,22 @@ async def test_can_have_rfr_role(num_roles, num_required, rfr_cog): else: assert rfr_cog.can_have_rfr_role(member) == any( x in required for x in member.roles), f"\n\r{member.roles}\n\r{required}" + + +@pytest.mark.asyncio +async def test_get_first_emoji_from_str(): + config: dpytest.RunnerConfig = dpytest.get_config() + guild: discord.Guild = config.guilds[0] + channel: discord.TextChannel = guild.text_channels[0] + + message: discord.Message = await dpytest.message("rfr") + msg_id: int = message.id + add_rfr_message(guild.id, channel.id, msg_id) + + emoji: discord.Emoji = testutils.fake_guild_emoji(guild) + role: discord.Role = testutils.fake_guild_role(guild) + + assert core.get_first_emoji_from_str(koalabot, guild, emoji) + + uni_emoji = testutils.fake_unicode_emoji() + assert core.get_first_emoji_from_str(koalabot, guild, uni_emoji) diff --git a/tests/cogs/react_for_role/test_db.py b/tests/cogs/react_for_role/test_db.py index 816ce90b..77fd5b06 100644 --- a/tests/cogs/react_for_role/test_db.py +++ b/tests/cogs/react_for_role/test_db.py @@ -20,12 +20,12 @@ from discord.ext.test import factories as dpyfactory # Own modules -from koala.cogs.react_for_role.db import ReactForRoleDBManager from koala.db import session_manager from tests.tests_utils import utils as testutils from tests.log import logger -from .utils import DBManager, independent_get_guild_rfr_message, independent_get_rfr_message_emoji_role, \ +from koala.cogs.react_for_role.db import * +from .utils import independent_get_guild_rfr_message, independent_get_rfr_message_emoji_role, \ independent_get_guild_rfr_required_role, get_rfr_reaction_role_by_role_id @@ -46,7 +46,7 @@ async def test_rfr_db_functions_guild_rfr_messages(): session, guild.id, channel.id, msg_id) == expected_full_list assert independent_get_guild_rfr_message(session) == expected_full_list # Test on adding first message, 1 message, 1 channel, 1 guild - DBManager.add_rfr_message(guild.id, channel.id, msg_id) + add_rfr_message(guild.id, channel.id, msg_id) expected_full_list.append((guild.id, channel.id, msg_id, 1)) assert independent_get_guild_rfr_message(session) == expected_full_list assert independent_get_guild_rfr_message(session, guild.id, channel.id, msg_id) == [ @@ -57,36 +57,36 @@ async def test_rfr_db_functions_guild_rfr_messages(): "TestGuild2Channel1", guild2) msg_id = dpyfactory.make_id() dpytest.get_config().guilds.append(guild2) - DBManager.add_rfr_message(guild2.id, channel2.id, msg_id) + add_rfr_message(guild2.id, channel2.id, msg_id) expected_full_list.append((guild2.id, channel2.id, msg_id, 2)) assert independent_get_guild_rfr_message(session, guild2.id, channel2.id, msg_id) == [ expected_full_list[1]] assert independent_get_guild_rfr_message(session, guild2.id, channel2.id, msg_id)[ - 0] == DBManager.get_rfr_message(guild2.id, - channel2.id, - msg_id) + 0] == get_rfr_message(guild2.id, + channel2.id, + msg_id) assert independent_get_guild_rfr_message(session) == expected_full_list # 1 guild, 2 channels with 1 message each guild1channel2: discord.TextChannel = dpytest.back.make_text_channel( "TestGuild1Channel2", guild) msg_id = dpyfactory.make_id() - DBManager.add_rfr_message(guild.id, guild1channel2.id, msg_id) + add_rfr_message(guild.id, guild1channel2.id, msg_id) expected_full_list.append((guild.id, guild1channel2.id, msg_id, 3)) assert independent_get_guild_rfr_message( session, guild.id, guild1channel2.id, msg_id) == [expected_full_list[2]] assert independent_get_guild_rfr_message(session, guild.id, guild1channel2.id, msg_id)[ - 0] == DBManager.get_rfr_message( + 0] == get_rfr_message( guild.id, guild1channel2.id, msg_id) assert independent_get_guild_rfr_message(session) == expected_full_list assert independent_get_guild_rfr_message(session, guild.id) == [expected_full_list[0], expected_full_list[2]] # 1 guild, 1 channel, with 2 messages msg_id = dpyfactory.make_id() - DBManager.add_rfr_message(guild.id, channel.id, msg_id) + add_rfr_message(guild.id, channel.id, msg_id) expected_full_list.append((guild.id, channel.id, msg_id, 4)) assert independent_get_guild_rfr_message(session, guild.id, channel.id, msg_id) == [ expected_full_list[3]] - assert independent_get_guild_rfr_message(session, guild.id, channel.id, msg_id)[0] == DBManager.get_rfr_message( + assert independent_get_guild_rfr_message(session, guild.id, channel.id, msg_id)[0] == get_rfr_message( guild.id, channel.id, msg_id) @@ -97,7 +97,7 @@ async def test_rfr_db_functions_guild_rfr_messages(): guild_rfr_messages = independent_get_guild_rfr_message(session) for guild_rfr_message in guild_rfr_messages: assert guild_rfr_message in guild_rfr_messages - DBManager.remove_rfr_message( + remove_rfr_message( guild_rfr_message[0], guild_rfr_message[1], guild_rfr_message[2]) assert guild_rfr_message not in independent_get_guild_rfr_message(session) assert independent_get_guild_rfr_message(session) == [] @@ -109,7 +109,7 @@ async def test_rfr_db_functions_rfr_message_emoji_roles(): guild: discord.Guild = dpytest.get_config().guilds[0] channel: discord.TextChannel = dpytest.get_config().channels[0] msg_id = dpyfactory.make_id() - DBManager.add_rfr_message(guild.id, channel.id, msg_id) + add_rfr_message(guild.id, channel.id, msg_id) guild_rfr_message = independent_get_guild_rfr_message(session)[0] expected_full_list: List[Tuple[int, str, int]] = [] assert independent_get_rfr_message_emoji_role(session) == expected_full_list @@ -117,113 +117,113 @@ async def test_rfr_db_functions_rfr_message_emoji_roles(): fake_emoji_1 = testutils.fake_unicode_emoji() fake_role_id_1 = dpyfactory.make_id() expected_full_list.append((1, fake_emoji_1, fake_role_id_1)) - DBManager.add_rfr_message_emoji_role( + add_rfr_message_emoji_role( guild_rfr_message[3], fake_emoji_1, fake_role_id_1) assert independent_get_rfr_message_emoji_role( - session) == expected_full_list, DBManager.get_rfr_message_emoji_roles(1) + session) == expected_full_list, get_rfr_message_emoji_roles(1) assert independent_get_rfr_message_emoji_role(session, 1) == expected_full_list assert independent_get_rfr_message_emoji_role(session, guild_rfr_message[3], fake_emoji_1, - fake_role_id_1) == [DBManager.get_rfr_reaction_role( + fake_role_id_1) == [get_rfr_reaction_role( guild_rfr_message[3], fake_emoji_1, fake_role_id_1)] # 1 unicode, 1 custom, trying to get same role fake_emoji_2 = testutils.fake_custom_emoji_str_rep() - DBManager.add_rfr_message_emoji_role( + add_rfr_message_emoji_role( guild_rfr_message[3], fake_emoji_2, fake_role_id_1) assert independent_get_rfr_message_emoji_role(session) == expected_full_list assert independent_get_rfr_message_emoji_role(session, - guild_rfr_message[3]) == DBManager.get_rfr_message_emoji_roles( + guild_rfr_message[3]) == get_rfr_message_emoji_roles( guild_rfr_message[3]) - assert [DBManager.get_rfr_reaction_role( + assert [get_rfr_reaction_role( guild_rfr_message[3], fake_emoji_2, fake_role_id_1)] == [None] # 2 roles, with 1 emoji trying to give both roles fake_role_id_2 = dpyfactory.make_id() - DBManager.add_rfr_message_emoji_role( + add_rfr_message_emoji_role( guild_rfr_message[3], fake_emoji_1, fake_role_id_2) assert independent_get_rfr_message_emoji_role(session) == expected_full_list assert independent_get_rfr_message_emoji_role(session, - guild_rfr_message[3]) == DBManager.get_rfr_message_emoji_roles( + guild_rfr_message[3]) == get_rfr_message_emoji_roles( guild_rfr_message[3]) - assert [DBManager.get_rfr_reaction_role( + assert [get_rfr_reaction_role( guild_rfr_message[3], fake_emoji_1, fake_role_id_2)] == [None] # 2 roles, 2 emojis, 1 message. split between them fake_emoji_2 = testutils.fake_custom_emoji_str_rep() fake_role_id_2 = dpyfactory.make_id() expected_full_list.append((1, fake_emoji_2, fake_role_id_2)) - DBManager.add_rfr_message_emoji_role(*expected_full_list[1]) + add_rfr_message_emoji_role(*expected_full_list[1]) assert independent_get_rfr_message_emoji_role(session) == expected_full_list assert independent_get_rfr_message_emoji_role(session, 1, fake_emoji_1) == [(1, fake_emoji_1, fake_role_id_1)] assert independent_get_rfr_message_emoji_role(session, 1, fake_emoji_2) == [(1, fake_emoji_2, fake_role_id_2)] assert independent_get_rfr_message_emoji_role(session, 1, fake_emoji_1)[0][ - 2] == DBManager.get_rfr_reaction_role_by_emoji_str(1, - fake_emoji_1) + 2] == get_rfr_reaction_role_by_emoji_str(1, + fake_emoji_1) assert independent_get_rfr_message_emoji_role(session, - 1) == DBManager.get_rfr_message_emoji_roles(1) + 1) == get_rfr_message_emoji_roles(1) assert independent_get_rfr_message_emoji_role(session, 1, role_id=fake_role_id_2)[0][ 2] == get_rfr_reaction_role_by_role_id(session, emoji_role_id=1, role_id=fake_role_id_2) # 2 roles 2 emojis, 2 messages. duplicated messages msg2_id = dpyfactory.make_id() - DBManager.add_rfr_message(guild.id, channel.id, msg2_id) + add_rfr_message(guild.id, channel.id, msg2_id) assert independent_get_guild_rfr_message(session ) == [guild_rfr_message, (guild.id, channel.id, msg2_id, 2)] guild_rfr_message_2 = independent_get_guild_rfr_message(session)[1] - DBManager.add_rfr_message_emoji_role( + add_rfr_message_emoji_role( guild_rfr_message_2[3], fake_emoji_1, fake_role_id_1) - DBManager.add_rfr_message_emoji_role( + add_rfr_message_emoji_role( guild_rfr_message_2[3], fake_emoji_2, fake_role_id_2) expected_full_list.extend([(guild_rfr_message_2[3], fake_emoji_1, fake_role_id_1), (guild_rfr_message_2[3], fake_emoji_2, fake_role_id_2)]) assert independent_get_rfr_message_emoji_role(session) == expected_full_list assert independent_get_rfr_message_emoji_role(session, - 2) == DBManager.get_rfr_message_emoji_roles(2) + 2) == get_rfr_message_emoji_roles(2) assert independent_get_rfr_message_emoji_role(session, - 1) == DBManager.get_rfr_message_emoji_roles(1) + 1) == get_rfr_message_emoji_roles(1) # 2 roles 2 emojis 2 messages. Swapped msg3_id = dpyfactory.make_id() - DBManager.add_rfr_message(guild.id, channel.id, msg3_id) + add_rfr_message(guild.id, channel.id, msg3_id) assert independent_get_guild_rfr_message(session) == [guild_rfr_message, (guild.id, channel.id, msg2_id, 2), (guild.id, channel.id, msg3_id, 3)] guild_rfr_message_3 = independent_get_guild_rfr_message(session)[2] - DBManager.add_rfr_message_emoji_role( + add_rfr_message_emoji_role( guild_rfr_message_3[3], fake_emoji_1, fake_role_id_2) - DBManager.add_rfr_message_emoji_role( + add_rfr_message_emoji_role( guild_rfr_message_3[3], fake_emoji_2, fake_role_id_1) expected_full_list.extend([(guild_rfr_message_3[3], fake_emoji_1, fake_role_id_2), (guild_rfr_message_3[3], fake_emoji_2, fake_role_id_1)]) assert independent_get_rfr_message_emoji_role(session) == expected_full_list assert independent_get_rfr_message_emoji_role(session, - 3) == DBManager.get_rfr_message_emoji_roles(3) + 3) == get_rfr_message_emoji_roles(3) assert [x[2] for x in independent_get_rfr_message_emoji_role(session, emoji_raw=fake_emoji_1)] == [ - DBManager.get_rfr_reaction_role_by_emoji_str(1, fake_emoji_1), - DBManager.get_rfr_reaction_role_by_emoji_str(2, fake_emoji_1), - DBManager.get_rfr_reaction_role_by_emoji_str(3, fake_emoji_1)] + get_rfr_reaction_role_by_emoji_str(1, fake_emoji_1), + get_rfr_reaction_role_by_emoji_str(2, fake_emoji_1), + get_rfr_reaction_role_by_emoji_str(3, fake_emoji_1)] assert [x[2] for x in independent_get_rfr_message_emoji_role(session, emoji_raw=fake_emoji_2)] == [ - DBManager.get_rfr_reaction_role_by_emoji_str(1, fake_emoji_2), - DBManager.get_rfr_reaction_role_by_emoji_str(2, fake_emoji_2), - DBManager.get_rfr_reaction_role_by_emoji_str(3, fake_emoji_2)] + get_rfr_reaction_role_by_emoji_str(1, fake_emoji_2), + get_rfr_reaction_role_by_emoji_str(2, fake_emoji_2), + get_rfr_reaction_role_by_emoji_str(3, fake_emoji_2)] # test deletion works from rfr message rfr_message_emoji_roles = independent_get_rfr_message_emoji_role(session, 3) - DBManager.remove_rfr_message(guild.id, channel.id, msg3_id) + remove_rfr_message(guild.id, channel.id, msg3_id) for row in rfr_message_emoji_roles: assert row not in independent_get_rfr_message_emoji_role(session ), independent_get_guild_rfr_message(session) # test deleting just emoji role combos rfr_message_emoji_roles = independent_get_rfr_message_emoji_role(session, 2) - DBManager.remove_rfr_message_emoji_roles(2) + remove_rfr_message_emoji_roles(2) for row in rfr_message_emoji_roles: assert row not in independent_get_rfr_message_emoji_role(session ), independent_get_guild_rfr_message(session) # test deleteing specific rfr_message_emoji_roles = independent_get_rfr_message_emoji_role(session, 1) - DBManager.remove_rfr_message_emoji_role( + remove_rfr_message_emoji_role( 1, emoji_raw=rfr_message_emoji_roles[0][1]) assert (rfr_message_emoji_roles[0][0], rfr_message_emoji_roles[0][1], rfr_message_emoji_roles[0][2]) not in independent_get_rfr_message_emoji_role(session) - DBManager.remove_rfr_message_emoji_role( + remove_rfr_message_emoji_role( 1, role_id=rfr_message_emoji_roles[1][2]) assert (rfr_message_emoji_roles[1][0], rfr_message_emoji_roles[1][1], rfr_message_emoji_roles[1][2]) not in independent_get_rfr_message_emoji_role(session) @@ -237,18 +237,18 @@ async def test_rfr_db_functions_guild_rfr_required_roles(): for i in range(50): role: discord.Role = testutils.fake_guild_role(guild) roles.append(role) - DBManager.add_guild_rfr_required_role(guild.id, role.id) + add_guild_rfr_required_role(guild.id, role.id) assert [x[1] for x in independent_get_guild_rfr_required_role(session)] == [x.id for x in roles], i assert [x[1] for x in - independent_get_guild_rfr_required_role(session)] == DBManager.get_guild_rfr_required_roles( + independent_get_guild_rfr_required_role(session)] == get_guild_rfr_required_roles( guild.id), i while len(roles) > 0: role: discord.Role = roles.pop() - DBManager.remove_guild_rfr_required_role(guild.id, role.id) + remove_guild_rfr_required_role(guild.id, role.id) assert [x[1] for x in independent_get_guild_rfr_required_role(session)] == [x.id for x in roles], len(roles) assert [x[1] for x in - independent_get_guild_rfr_required_role(session)] == DBManager.get_guild_rfr_required_roles( + independent_get_guild_rfr_required_role(session)] == get_guild_rfr_required_roles( guild.id), len(roles) @@ -266,7 +266,7 @@ async def test_rfr_without_req_role(num_roles, num_required, rfr_cog): r_list.append(role) required = random.sample(list(r_list), num_required) for r in required: - DBManager.add_guild_rfr_required_role(test_guild.id, r.id) + add_guild_rfr_required_role(test_guild.id, r.id) assert independent_get_guild_rfr_required_role(session, test_guild.id, r.id) is not None member: discord.Member = await dpytest.member_join() @@ -276,13 +276,13 @@ async def test_rfr_without_req_role(num_roles, num_required, rfr_cog): # Create RFR message for test rfr_message = dpytest.back.make_message("FakeContent", config.client.user, test_guild.text_channels[0]) - DBManager.add_rfr_message(test_guild.id, rfr_message.channel.id, rfr_message.id) - assert DBManager.get_rfr_message(test_guild.id, rfr_message.channel.id, rfr_message.id) is not None + add_rfr_message(test_guild.id, rfr_message.channel.id, rfr_message.id) + assert get_rfr_message(test_guild.id, rfr_message.channel.id, rfr_message.id) is not None # Add emoji role combo to db - _, _, _, er_id = DBManager.get_rfr_message(test_guild.id, rfr_message.channel.id, rfr_message.id) + _, _, _, er_id = get_rfr_message(test_guild.id, rfr_message.channel.id, rfr_message.id) react_emoji: str = testutils.fake_unicode_emoji() - DBManager.add_rfr_message_emoji_role(er_id, emoji.demojize(react_emoji), role_to_add.id) + add_rfr_message_emoji_role(er_id, emoji.demojize(react_emoji), role_to_add.id) with mock.patch("koala.cogs.ReactForRole.get_role_member_info", mock.AsyncMock(return_value=(member, role_to_add))): @@ -303,8 +303,8 @@ async def test_rfr_with_req_role(num_roles, num_required, rfr_cog): # Create RFR message for test rfr_message = dpytest.back.make_message("FakeContent", config.client.user, test_guild.text_channels[0]) - DBManager.add_rfr_message(test_guild.id, rfr_message.channel.id, rfr_message.id) - assert DBManager.get_rfr_message(test_guild.id, rfr_message.channel.id, rfr_message.id) is not None + add_rfr_message(test_guild.id, rfr_message.channel.id, rfr_message.id) + assert get_rfr_message(test_guild.id, rfr_message.channel.id, rfr_message.id) is not None r_list = [] for i in range(num_roles): @@ -314,12 +314,12 @@ async def test_rfr_with_req_role(num_roles, num_required, rfr_cog): role_to_add = testutils.fake_guild_role(test_guild) # Add emoji role combo to db - _, _, _, er_id = DBManager.get_rfr_message(test_guild.id, rfr_message.channel.id, rfr_message.id) + _, _, _, er_id = get_rfr_message(test_guild.id, rfr_message.channel.id, rfr_message.id) react_emoji: str = testutils.fake_unicode_emoji() - DBManager.add_rfr_message_emoji_role(er_id, emoji.demojize(react_emoji), role_to_add.id) + add_rfr_message_emoji_role(er_id, emoji.demojize(react_emoji), role_to_add.id) for r in required: - DBManager.add_guild_rfr_required_role(test_guild.id, r.id) + add_guild_rfr_required_role(test_guild.id, r.id) assert independent_get_guild_rfr_required_role(session, test_guild.id, r.id) is not None member: discord.Member = await dpytest.member_join() diff --git a/tests/cogs/react_for_role/utils.py b/tests/cogs/react_for_role/utils.py index 90f33fb9..b3a0d259 100644 --- a/tests/cogs/react_for_role/utils.py +++ b/tests/cogs/react_for_role/utils.py @@ -15,14 +15,8 @@ from sqlalchemy import select # Own modules -from koala.cogs.react_for_role.db import ReactForRoleDBManager from koala.cogs.react_for_role.models import GuildRFRRequiredRoles, GuildRFRMessages, RFRMessageEmojiRoles -# Constants - -# Variables -DBManager = ReactForRoleDBManager() - def independent_get_guild_rfr_message(session: sqlalchemy.orm.Session, guild_id=None, channel_id=None, message_id=None ) -> List[Tuple[int, int, int, int]]: @@ -37,7 +31,8 @@ def independent_get_guild_rfr_message(session: sqlalchemy.orm.Session, guild_id= return [row.old_format() for row in rows] -def independent_get_rfr_message_emoji_role(session: sqlalchemy.orm.Session, emoji_role_id=None, emoji_raw=None, role_id=None) -> List[ +def independent_get_rfr_message_emoji_role(session: sqlalchemy.orm.Session, emoji_role_id=None, emoji_raw=None, + role_id=None) -> List[ Tuple[int, str, int]]: sql_select = select(RFRMessageEmojiRoles) if emoji_role_id is not None: