From ed578f383c09711bcd3aea7b1c6270d28062783e Mon Sep 17 00:00:00 2001 From: Stefan Cooper Date: Fri, 20 May 2022 17:20:42 +0100 Subject: [PATCH 01/25] feat: add rfr core/db Contributes to: KoalaBotUK/KoalaBotFrontend#439 Signed-off-by: Stefan Cooper --- koala/cogs/react_for_role/__init__.py | 5 +- koala/cogs/react_for_role/cog.py | 47 ++--- koala/cogs/react_for_role/core.py | 67 +++++++ koala/cogs/react_for_role/db2.py | 267 ++++++++++++++++++++++++++ tests/cogs/react_for_role/test_cog.py | 6 +- 5 files changed, 368 insertions(+), 24 deletions(-) create mode 100644 koala/cogs/react_for_role/core.py create mode 100644 koala/cogs/react_for_role/db2.py diff --git a/koala/cogs/react_for_role/__init__.py b/koala/cogs/react_for_role/__init__.py index f38c35dd..62419b41 100644 --- a/koala/cogs/react_for_role/__init__.py +++ b/koala/cogs/react_for_role/__init__.py @@ -1,2 +1,5 @@ -from . import utils, db, models +from . import utils, db, models, cog, core from .cog import ReactForRole, setup + +def setup(bot): + cog.setup(bot) \ No newline at end of file diff --git a/koala/cogs/react_for_role/cog.py b/koala/cogs/react_for_role/cog.py index 15944628..005fe8a0 100644 --- a/koala/cogs/react_for_role/cog.py +++ b/koala/cogs/react_for_role/cog.py @@ -18,6 +18,7 @@ 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 @@ -181,12 +182,10 @@ 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 = await core.create_rfr_message(title, ctx.guild, desc, KOALA_GREEN, channel) + # 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 " @@ -398,7 +397,7 @@ 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) + await core.setup_rfr_reaction_permissions(chnl.guild, chnl, self.bot) emb = self.get_embed_from_message(msg) reacts: List[Union[discord.PartialEmoji, discord.Emoji, str]] = [x.emoji for x in msg.reactions] if not emb: @@ -905,6 +904,7 @@ async def parse_emoji_or_roles_input_str(self, ctx: commands.Context, input_str: else: arr.append(raw_emoji) return arr + async def prompt_for_input(self, ctx: commands.Context, input_type: str) -> Union[discord.Attachment, str]: """ @@ -943,6 +943,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,21 +955,23 @@ 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_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 + # """ + # print("BBBBBBBBBBB") + # print(msg.embeds) + # 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: """ diff --git a/koala/cogs/react_for_role/core.py b/koala/cogs/react_for_role/core.py new file mode 100644 index 00000000..aea9bc48 --- /dev/null +++ b/koala/cogs/react_for_role/core.py @@ -0,0 +1,67 @@ +import datetime +from typing import List, Optional + +import discord +from discord.ext.commands import Bot + +from . import db2 +from .log import logger + +from koala.db import assign_session +import discord +from discord import Colour +# Constants + +koala_logo = "https://cdn.discordapp.com/attachments/737280260541907015/752024535985029240/discord1.png" + +# Variables +# current_activity = None + +@assign_session +async def create_rfr_message(title: str, guild: discord.Guild, description: str, colour: Colour, channel: discord.TextChannel, **kwargs): + embed: discord.Embed = discord.Embed(title=title, description=description, colour=colour) + embed.set_footer(text="ReactForRole") + embed.set_thumbnail(url=koala_logo) + rfr_msg: discord.Message = await channel.send(embed=embed) + db2.add_rfr_message(guild.id, channel.id, rfr_msg.id, **kwargs) + return rfr_msg + + +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 + """ + + # TODO: Figure out a way to get this working in core + + if not msg: + return None + try: + embed = msg.embeds[0] + if not embed: + return None + return embed + except IndexError: + return None diff --git a/koala/cogs/react_for_role/db2.py b/koala/cogs/react_for_role/db2.py new file mode 100644 index 00000000..fdd79801 --- /dev/null +++ b/koala/cogs/react_for_role/db2.py @@ -0,0 +1,267 @@ +#!/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 + +# Libs + +# Constants + + +# class ReactForRoleDBManager: +# """ +# A class for interacting with the KoalaBot ReactForRole database +# """ + +@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): + """ + 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) + +@assign_session +def remove_rfr_message_emoji_role(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() + +@assign_session +def remove_rfr_message_emoji_roles(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() + +@assign_session +def remove_rfr_message(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: + 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() + +@assign_session +def get_rfr_message(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 + +@assign_session +def get_guild_rfr_messages(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] + +@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]] = 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 + +@assign_session +def get_rfr_message_emoji_roles(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] + +@assign_session +def get_rfr_reaction_role(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 + +@assign_session +def get_rfr_reaction_role_by_emoji_str(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] + +@assign_session +def add_guild_rfr_required_role(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() + +@assign_session +def remove_guild_rfr_required_role(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() + +@assign_session +def get_guild_rfr_required_roles(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 diff --git a/tests/cogs/react_for_role/test_cog.py b/tests/cogs/react_for_role/test_cog.py index df7de959..41bf6cba 100644 --- a/tests/cogs/react_for_role/test_cog.py +++ b/tests/cogs/react_for_role/test_cog.py @@ -20,12 +20,14 @@ from discord.ext.test import factories as dpyfactory # Own modules +from koala.cogs.react_for_role import core import koalabot 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 tests.log import logger +from koala.cogs import ReactForRole # Constants @@ -195,14 +197,16 @@ async def test_prompt_for_input_attachment(rfr_cog, utils_cog): @pytest.mark.asyncio -async def test_overwrite_channel_add_reaction_perms(rfr_cog): +async def test_setup_rfr_reaction_permissions(rfr_cog): 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, 'role', reason=None), mock.call(channel.id, config.client.user.id, 64, 0, 'member', From ed259c1e2968e9b60a577922119f469d0861b365 Mon Sep 17 00:00:00 2001 From: JayDwee Date: Wed, 22 Jun 2022 19:44:43 +0100 Subject: [PATCH 02/25] test: fix get_embed test --- tests/cogs/react_for_role/test_cog.py | 13 ++++++------- 1 file changed, 6 insertions(+), 7 deletions(-) diff --git a/tests/cogs/react_for_role/test_cog.py b/tests/cogs/react_for_role/test_cog.py index 41bf6cba..1a33272b 100644 --- a/tests/cogs/react_for_role/test_cog.py +++ b/tests/cogs/react_for_role/test_cog.py @@ -229,21 +229,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, '', embed=test_embed_dict) + 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]) From 4a19ef6757dde327d4cb6803add1297d730512cf Mon Sep 17 00:00:00 2001 From: Stefan Cooper Date: Mon, 27 Jun 2022 18:30:55 +0100 Subject: [PATCH 03/25] fix: move cog to core Contributes to: KoalaBotUK/KoalaBotFrontend#439 Signed-off-by: Stefan Cooper --- koala/cogs/react_for_role/cog.py | 225 +++++-------------------- koala/cogs/react_for_role/core.py | 229 ++++++++++++++++++++++---- koala/cogs/react_for_role/db2.py | 132 ++++++--------- tests/cogs/react_for_role/test_cog.py | 22 +-- 4 files changed, 308 insertions(+), 300 deletions(-) diff --git a/koala/cogs/react_for_role/cog.py b/koala/cogs/react_for_role/cog.py index 005fe8a0..2e72a34f 100644 --- a/koala/cogs/react_for_role/cog.py +++ b/koala/cogs/react_for_role/cog.py @@ -25,7 +25,6 @@ from koala.db import insert_extension from .db import ReactForRoleDBManager from .log import logger -from .utils import CUSTOM_EMOJI_REGEXP, UNICODE_EMOJI_REGEXP def rfr_is_enabled(ctx): @@ -48,7 +47,7 @@ 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() @@ -208,10 +207,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(ctx.guild.id, channel.id, msg) await ctx.send("ReactForRole Message deleted") else: await ctx.send("Cancelled command.") @@ -233,14 +229,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(embed, msg, description=desc) else: await ctx.send("Okay, cancelling command.") else: @@ -259,14 +254,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(embed, msg, title=title) else: await ctx.send("Okay, cancelling command.") else: @@ -285,7 +279,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}.") @@ -299,15 +293,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(embed, msg, image_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(embed, msg, image_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.") @@ -350,25 +342,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(id=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: @@ -382,11 +362,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(embed, msg) await ctx.send("Okay, should be done. Please check.") @commands.check(koalabot.is_admin) @@ -398,7 +374,7 @@ async def rfr_fix_embed(self, ctx: commands.Context): """ msg, chnl = await self.get_rfr_message_from_prompts(ctx) await core.setup_rfr_reaction_permissions(chnl.guild, chnl, self.bot) - emb = self.get_embed_from_message(msg) + 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( @@ -457,7 +433,7 @@ async def rfr_add_roles_to_msg(self, ctx: commands.Context): 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( @@ -468,14 +444,9 @@ 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) + old_embed = core.get_embed_from_message(msg) + msg = core.create_rfr_message(title=old_embed.title, guild=ctx.guild, description=old_embed.description, colour=KOALA_GREEN, channel=channel) 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) else: @@ -491,35 +462,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) + rfr_embed = core.get_embed_from_message(msg) + duplicateRolesFound, duplicateEmojisFound, edited_msg = core.rfr_add_emoji_role(ctx.guild, channel, rfr_embed, msg, rfr_msg_row, 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) @@ -546,7 +492,7 @@ async def rfr_remove_roles_from_msg(self, ctx: commands.Context): 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( @@ -555,8 +501,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(ctx.guild.id, channel.id, msg) await ctx.send("Okay, deleted that react for role message. Have a nice day.") return else: @@ -570,54 +515,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_fields = [field for field in rfr_embed_fields if - rfr_embed_fields.index(field) not in removed_field_indexes] + 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) - 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(ctx.guild.id, channel.id, msg) 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() @@ -687,9 +598,8 @@ async def rfr_add_guild_required_role(self, ctx: commands.Context, role_str: str :return: """ try: - role: discord.Role = await commands.RoleConverter().convert(ctx, role_str) + role: discord.Role = await core.add_guild_rfr_required_role(self.bot, ctx.guild, role_str) 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.") @@ -706,10 +616,9 @@ async def rfr_remove_guild_required_role(self, ctx: commands.Context, role_str: :return: """ try: - role: discord.Role = await commands.RoleConverter().convert(ctx, role_str) + role: discord.Role = await core.remove_guild_rfr_required_role(self.bot, ctx.guild, role_str) 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.") @@ -723,7 +632,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.id) msg_str = "You will need one of these roles to react to rfr messages on this server:\n" for role_id in role_ids: @@ -819,7 +728,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) @@ -862,15 +771,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: @@ -894,7 +809,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: @@ -955,62 +872,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 - # """ - # print("BBBBBBBBBBB") - # print(msg.embeds) - # 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 index aea9bc48..9fcbfe82 100644 --- a/koala/cogs/react_for_role/core.py +++ b/koala/cogs/react_for_role/core.py @@ -1,8 +1,11 @@ +from ast import Tuple import datetime -from typing import List, Optional +from typing import * import discord from discord.ext.commands import Bot +from discord.ext import commands +import emoji from . import db2 from .log import logger @@ -10,6 +13,8 @@ from koala.db import assign_session import discord from discord import Colour +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" @@ -17,35 +22,167 @@ # Variables # current_activity = None +def create_ctx(bot: Bot, guild: discord.Guild): + return { 'bot': bot, 'guild': guild } + @assign_session async def create_rfr_message(title: str, guild: discord.Guild, description: str, colour: Colour, channel: discord.TextChannel, **kwargs): - embed: discord.Embed = discord.Embed(title=title, description=description, colour=colour) - embed.set_footer(text="ReactForRole") - embed.set_thumbnail(url=koala_logo) - rfr_msg: discord.Message = await channel.send(embed=embed) - db2.add_rfr_message(guild.id, channel.id, rfr_msg.id, **kwargs) - return rfr_msg + embed: discord.Embed = discord.Embed(title=title, description=description, colour=colour) + embed.set_footer(text="ReactForRole") + embed.set_thumbnail(url=koala_logo) + rfr_msg: discord.Message = await channel.send(embed=embed) + db2.add_rfr_message(guild.id, channel.id, rfr_msg.id, **kwargs) + return rfr_msg + +@assign_session +async def delete_rfr_message(guild_id: str, channel_id: str, msg: discord.Message, **kwargs): + rfr_msg_row = db2.get_rfr_message(guild_id, channel_id, msg.id, **kwargs) + db2.remove_rfr_message_emoji_roles(rfr_msg_row[3], **kwargs) + db2.remove_rfr_message(guild_id, channel_id, msg.id, **kwargs) + await msg.delete() + +@assign_session +async def use_inline_rfr_all(guild: discord.Guild, **kwargs): + text_channels: List[discord.TextChannel] = guild.text_channels + guild_rfr_messages = db2.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(id=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(embed: discord.Embed, msg: discord.Message): + 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 rfr_edit(embed: discord.Embed, msg: discord.Message, description: str = "", title: str = "", image_url: str = ""): + embed.description = description + embed.title = title + embed.set_thumbnail(url=image_url) + await msg.edit(embed=embed) + return msg + +@assign_session +async def rfr_remove_emojis_roles(bot: Bot, guild: discord.Guild, msg: discord.Message, rfr_msg_row: discord.Message, 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 = discord.Embed(title=rfr_embed.title, description=rfr_embed.description, + colour=KOALA_GREEN) + new_embed.set_thumbnail( + url=koala_logo) + new_embed.set_footer(text="ReactForRole") + 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): + db2.remove_rfr_message_emoji_role(rfr_msg_row[3], emoji_raw=emoji.demojize(row), **kwargs) + else: + db2.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) + db2.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 = 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_add_emoji_role(guild: str, channel: discord.TextChannel, rfr_embed: discord.Embed, msg: discord.Message, rfr_msg_row: discord.Message, emoji_role_map: List[Tuple[Union[discord.Emoji, str], discord.Role]], **kwargs): + duplicateRolesFound = False + duplicateEmojisFound = False + + 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]: + duplicateEmojisFound = True + elif role in [x.value for x in rfr_embed.fields]: + duplicateRolesFound = True + else: + if isinstance(discord_emoji, str): + db2.add_rfr_message_emoji_role(rfr_msg_row[3], emoji.demojize(discord_emoji), + role.id, **kwargs) + else: + db2.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 {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 duplicateRolesFound, duplicateEmojisFound, edited_msg + +async def add_guild_rfr_required_role(bot: Bot, guild: discord.Guild, role_str: str, **kwargs): + ctx = create_ctx(bot, guild) + role: discord.Role = await commands.RoleConverter().convert(ctx, role_str) + db2.remove_guild_rfr_required_role(ctx.guild.id, role.id, **kwargs) + return role + +async def remove_guild_rfr_required_role(bot: Bot, guild: discord.Guild, role_str: str, **kwargs): + ctx = create_ctx(bot, guild) + role: discord.Role = await commands.RoleConverter().convert(ctx, role_str) + db2.add_guild_rfr_required_role(ctx.guild.id, role.id, **kwargs) + return role + +def rfr_list_guild_required_roles(guild: discord.Guild, **kwargs): + return db2.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) + """ + 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]: """ @@ -53,9 +190,6 @@ def get_embed_from_message(msg: discord.Message) -> Optional[discord.Embed]: :param msg: Message to check :return: Returns the embed if there is one. If there isn't returns None """ - - # TODO: Figure out a way to get this working in core - if not msg: return None try: @@ -65,3 +199,42 @@ def get_embed_from_message(msg: discord.Message) -> Optional[discord.Embed]: 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) -> 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. + """ + + ctx = create_ctx(bot, guild) + + # 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_str = search_result.group().strip() + try: + discord_emoji: discord.Emoji = await commands.EmojiConverter().convert(ctx, emoji_str) + 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?" \ No newline at end of file diff --git a/koala/cogs/react_for_role/db2.py b/koala/cogs/react_for_role/db2.py index fdd79801..6dac6fc4 100644 --- a/koala/cogs/react_for_role/db2.py +++ b/koala/cogs/react_for_role/db2.py @@ -1,15 +1,8 @@ #!/usr/bin/env python -""" -KoalaBot Reaction Roles Code - -Author: Anan Venkatesh -Commented using reStructuredText (reST) -""" -# Futures - # Built-in/Generic Imports from typing import * +from requests import Session import sqlalchemy.exc import sqlalchemy.orm @@ -21,16 +14,6 @@ from .models import GuildRFRMessages, RFRMessageEmojiRoles, GuildRFRRequiredRoles from koala.db import assign_session -# Libs - -# Constants - - -# class ReactForRoleDBManager: -# """ -# A class for interacting with the KoalaBot ReactForRole database -# """ - @assign_session def add_rfr_message(guild_id: int, channel_id: int, message_id: int, session: sqlalchemy.orm.Session): """ @@ -46,7 +29,7 @@ def add_rfr_message(guild_id: int, channel_id: int, message_id: int, session: sq session.commit() @assign_session -def add_rfr_message_emoji_role(emoji_role_id: int, emoji_raw: str, role_id: int): +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 @@ -54,16 +37,15 @@ def add_rfr_message_emoji_role(emoji_role_id: int, emoji_raw: str, role_id: int) :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) + 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): +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 @@ -86,26 +68,24 @@ def remove_rfr_message_emoji_role(emoji_role_id: int, emoji_raw: str = None, rol RFRMessageEmojiRoles.emoji_role_id == emoji_role_id, RFRMessageEmojiRoles.emoji_raw == emoji_raw )) - with session_manager() as session: - session.execute(delete_sql) - session.commit() + session.execute(delete_sql) + session.commit() @assign_session -def remove_rfr_message_emoji_roles(emoji_role_id: int): +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: """ - with session_manager() as session: - delete_sql = delete(RFRMessageEmojiRoles) \ - .where(RFRMessageEmojiRoles.emoji_role_id == emoji_role_id) + delete_sql = delete(RFRMessageEmojiRoles) \ + .where(RFRMessageEmojiRoles.emoji_role_id == emoji_role_id) - session.execute(delete_sql) - session.commit() + session.execute(delete_sql) + session.commit() @assign_session -def remove_rfr_message(guild_id: int, channel_id: int, message_id: int): +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 @@ -113,23 +93,22 @@ def remove_rfr_message(guild_id: int, channel_id: int, message_id: int): :param message_id: Message ID of the rfr message :return: """ - emoji_role_id = self.get_rfr_message(guild_id, channel_id, message_id) + emoji_role_id = get_rfr_message(guild_id, channel_id, message_id) if not emoji_role_id: 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() + 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) -> Optional[Tuple[int, int, int, int]]: +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 @@ -137,28 +116,26 @@ def get_rfr_message(guild_id: int, channel_id: int, message_id: int) -> Optional :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 + 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): +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. """ - 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] + messages = session.execute(select(GuildRFRMessages) + .filter_by(guild_id=guild_id)).scalars().all() + return [message.old_format() + for message in messages] @assign_session def get_guild_rfr_roles(guild_id: int) -> List[int]: @@ -228,40 +205,37 @@ def get_rfr_reaction_role_by_emoji_str(emoji_role_id: int, emoji_raw: str) -> Op return row[0] @assign_session -def add_guild_rfr_required_role(guild_id: int, role_id: int): +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: """ - with session_manager() as session: - session.add(GuildRFRRequiredRoles(guild_id=guild_id, role_id=role_id)) - session.commit() + 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): +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: """ - with session_manager() as session: - session.execute(delete(GuildRFRRequiredRoles).filter_by(guild_id=guild_id, role_id=role_id)) - session.commit() + 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) -> List[int]: +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 """ - with session_manager() as session: - rows = session.execute(select(GuildRFRRequiredRoles).filter_by(guild_id=guild_id)).scalars().all() + 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 + role_ids = [x.role_id for x in rows] + if not role_ids: + return [] + return role_ids diff --git a/tests/cogs/react_for_role/test_cog.py b/tests/cogs/react_for_role/test_cog.py index 1a33272b..bdba553b 100644 --- a/tests/cogs/react_for_role/test_cog.py +++ b/tests/cogs/react_for_role/test_cog.py @@ -260,12 +260,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() @@ -282,7 +282,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 @@ -374,7 +374,7 @@ async def test_rfr_edit_description(): 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() @@ -397,7 +397,7 @@ async def test_rfr_edit_title(): 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() @@ -428,7 +428,7 @@ async def test_rfr_edit_thumbnail_attach(): 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" @@ -453,7 +453,7 @@ async def test_rfr_edit_thumbnail_bad_attach(attach): 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: @@ -482,7 +482,7 @@ async def test_rfr_edit_thumbnail_links(image_url): 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") @@ -511,7 +511,7 @@ async def test_rfr_edit_inline_all(arg): 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() @@ -550,7 +550,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: @@ -585,7 +585,7 @@ async def test_rfr_remove_roles_from_msg(): 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: From 29c3bf9c531ed4e411336c8ba60e0bc2e81cd19f Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 31 Oct 2022 19:51:40 +0000 Subject: [PATCH 04/25] Remove db2.py and move the code to db.py --- koala/cogs/react_for_role/core.py | 28 +- koala/cogs/react_for_role/db.py | 462 +++++++++++++++--------------- koala/cogs/react_for_role/db2.py | 241 ---------------- 3 files changed, 239 insertions(+), 492 deletions(-) delete mode 100644 koala/cogs/react_for_role/db2.py diff --git a/koala/cogs/react_for_role/core.py b/koala/cogs/react_for_role/core.py index 9fcbfe82..e41313d6 100644 --- a/koala/cogs/react_for_role/core.py +++ b/koala/cogs/react_for_role/core.py @@ -7,7 +7,7 @@ from discord.ext import commands import emoji -from . import db2 +from . import db from .log import logger from koala.db import assign_session @@ -31,20 +31,20 @@ async def create_rfr_message(title: str, guild: discord.Guild, description: str, embed.set_footer(text="ReactForRole") embed.set_thumbnail(url=koala_logo) rfr_msg: discord.Message = await channel.send(embed=embed) - db2.add_rfr_message(guild.id, channel.id, rfr_msg.id, **kwargs) + db.add_rfr_message(guild.id, channel.id, rfr_msg.id, **kwargs) return rfr_msg @assign_session async def delete_rfr_message(guild_id: str, channel_id: str, msg: discord.Message, **kwargs): - rfr_msg_row = db2.get_rfr_message(guild_id, channel_id, msg.id, **kwargs) - db2.remove_rfr_message_emoji_roles(rfr_msg_row[3], **kwargs) - db2.remove_rfr_message(guild_id, channel_id, msg.id, **kwargs) + rfr_msg_row = db.get_rfr_message(guild_id, channel_id, msg.id, **kwargs) + db.remove_rfr_message_emoji_roles(rfr_msg_row[3], **kwargs) + db.remove_rfr_message(guild_id, channel_id, msg.id, **kwargs) await msg.delete() @assign_session async def use_inline_rfr_all(guild: discord.Guild, **kwargs): text_channels: List[discord.TextChannel] = guild.text_channels - guild_rfr_messages = db2.get_guild_rfr_messages(guild.id, **kwargs) + 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(id=rfr_message[2]) @@ -86,13 +86,13 @@ async def rfr_remove_emojis_roles(bot: Bot, guild: discord.Guild, msg: discord.M 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): - db2.remove_rfr_message_emoji_role(rfr_msg_row[3], emoji_raw=emoji.demojize(row), **kwargs) + db.remove_rfr_message_emoji_role(rfr_msg_row[3], emoji_raw=emoji.demojize(row), **kwargs) else: - db2.remove_rfr_message_emoji_role(rfr_msg_row[3], emoji_raw=row, **kwargs) + 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) - db2.remove_rfr_message_emoji_role(rfr_msg_row[3], role_id=row.id, **kwargs) + 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) @@ -130,10 +130,10 @@ async def rfr_add_emoji_role(guild: str, channel: discord.TextChannel, rfr_embed duplicateRolesFound = True else: if isinstance(discord_emoji, str): - db2.add_rfr_message_emoji_role(rfr_msg_row[3], emoji.demojize(discord_emoji), + db.add_rfr_message_emoji_role(rfr_msg_row[3], emoji.demojize(discord_emoji), role.id, **kwargs) else: - db2.add_rfr_message_emoji_role(rfr_msg_row[3], str(discord_emoji), role.id, **kwargs) + 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) @@ -152,17 +152,17 @@ async def rfr_add_emoji_role(guild: str, channel: discord.TextChannel, rfr_embed async def add_guild_rfr_required_role(bot: Bot, guild: discord.Guild, role_str: str, **kwargs): ctx = create_ctx(bot, guild) role: discord.Role = await commands.RoleConverter().convert(ctx, role_str) - db2.remove_guild_rfr_required_role(ctx.guild.id, role.id, **kwargs) + db.remove_guild_rfr_required_role(ctx.guild.id, role.id, **kwargs) return role async def remove_guild_rfr_required_role(bot: Bot, guild: discord.Guild, role_str: str, **kwargs): ctx = create_ctx(bot, guild) role: discord.Role = await commands.RoleConverter().convert(ctx, role_str) - db2.add_guild_rfr_required_role(ctx.guild.id, role.id, **kwargs) + db.add_guild_rfr_required_role(ctx.guild.id, role.id, **kwargs) return role def rfr_list_guild_required_roles(guild: discord.Guild, **kwargs): - return db2.get_guild_rfr_required_roles(guild.id, **kwargs) + return db.get_guild_rfr_required_roles(guild.id, **kwargs) async def setup_rfr_reaction_permissions(guild: discord.Guild, channel: discord.TextChannel, bot: Bot): """ diff --git a/koala/cogs/react_for_role/db.py b/koala/cogs/react_for_role/db.py index df283dc1..52cc18d7 100644 --- a/koala/cogs/react_for_role/db.py +++ b/koala/cogs/react_for_role/db.py @@ -1,253 +1,241 @@ #!/usr/bin/env python -""" -KoalaBot Reaction Roles Code - -Author: Anan Venkatesh -Commented using reStructuredText (reST) -""" -# Futures - # Built-in/Generic Imports from typing import * +from requests import Session 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] + +@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]] = 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 + +@assign_session +def get_rfr_message_emoji_roles(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] -# 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 - )) +@assign_session +def get_rfr_reaction_role(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: - 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: + return None + +@assign_session +def get_rfr_reaction_role_by_emoji_str(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 - 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 \ No newline at end of file diff --git a/koala/cogs/react_for_role/db2.py b/koala/cogs/react_for_role/db2.py deleted file mode 100644 index 6dac6fc4..00000000 --- a/koala/cogs/react_for_role/db2.py +++ /dev/null @@ -1,241 +0,0 @@ -#!/usr/bin/env python - -# Built-in/Generic Imports -from typing import * -from requests import Session - -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] - -@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]] = 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 - -@assign_session -def get_rfr_message_emoji_roles(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] - -@assign_session -def get_rfr_reaction_role(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 - -@assign_session -def get_rfr_reaction_role_by_emoji_str(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] - -@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 From 5bb8a159a67a075d3fc15607532ae6b83093003a Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 31 Oct 2022 20:00:30 +0000 Subject: [PATCH 05/25] More re-factoring --- koala/cogs/react_for_role/db.py | 444 ++++++++++++++++---------------- 1 file changed, 223 insertions(+), 221 deletions(-) diff --git a/koala/cogs/react_for_role/db.py b/koala/cogs/react_for_role/db.py index 52cc18d7..2314c754 100644 --- a/koala/cogs/react_for_role/db.py +++ b/koala/cogs/react_for_role/db.py @@ -14,228 +14,230 @@ 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)) +class ReactForRoleDBManager: + + @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() - 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] - -@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]] = 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 - -@assign_session -def get_rfr_message_emoji_roles(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] - -@assign_session -def get_rfr_reaction_role(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 + + @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: - return None + 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 get_rfr_reaction_role_by_emoji_str(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: + @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 = ReactForRoleDBManager.get_rfr_message(guild_id, channel_id, message_id) + if not emoji_role_id: return - 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 \ No newline at end of file + else: + ReactForRoleDBManager.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] + + @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]] = ReactForRoleDBManager.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): + """ + 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] + + @assign_session + def get_rfr_reaction_role(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 + + @assign_session + def get_rfr_reaction_role_by_emoji_str(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] + + @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 \ No newline at end of file From 30626e4b9da79ac4411c87d13de9571a76109038 Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 31 Oct 2022 20:21:44 +0000 Subject: [PATCH 06/25] Fix double importing in core.py --- koala/cogs/react_for_role/core.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/koala/cogs/react_for_role/core.py b/koala/cogs/react_for_role/core.py index e41313d6..6d6f8eac 100644 --- a/koala/cogs/react_for_role/core.py +++ b/koala/cogs/react_for_role/core.py @@ -1,5 +1,4 @@ from ast import Tuple -import datetime from typing import * import discord @@ -11,8 +10,6 @@ from .log import logger from koala.db import assign_session -import discord -from discord import Colour from koala.colours import KOALA_GREEN from .utils import CUSTOM_EMOJI_REGEXP, UNICODE_EMOJI_REGEXP # Constants @@ -26,7 +23,7 @@ def create_ctx(bot: Bot, guild: discord.Guild): return { 'bot': bot, 'guild': guild } @assign_session -async def create_rfr_message(title: str, guild: discord.Guild, description: str, colour: Colour, channel: discord.TextChannel, **kwargs): +async def create_rfr_message(title: str, guild: discord.Guild, description: str, colour: discord.Colour, channel: discord.TextChannel, **kwargs): embed: discord.Embed = discord.Embed(title=title, description=description, colour=colour) embed.set_footer(text="ReactForRole") embed.set_thumbnail(url=koala_logo) From 1716a077433eac498f3d85e0e680946f9b0b29ac Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 31 Oct 2022 20:22:10 +0000 Subject: [PATCH 07/25] Fix referencing to db.py --- koala/cogs/react_for_role/cog.py | 25 +- koala/cogs/react_for_role/db.py | 444 +++++++++++++++---------------- 2 files changed, 233 insertions(+), 236 deletions(-) diff --git a/koala/cogs/react_for_role/cog.py b/koala/cogs/react_for_role/cog.py index 2e72a34f..01f404c5 100644 --- a/koala/cogs/react_for_role/cog.py +++ b/koala/cogs/react_for_role/cog.py @@ -23,7 +23,7 @@ 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 * from .log import logger @@ -50,7 +50,6 @@ class ReactForRole(commands.Cog): 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) @@ -380,12 +379,12 @@ async def rfr_fix_embed(self, ctx: commands.Context): 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.") @@ -428,7 +427,7 @@ 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.") @@ -448,7 +447,7 @@ async def rfr_add_roles_to_msg(self, ctx: commands.Context): msg = core.create_rfr_message(title=old_embed.title, guild=ctx.guild, description=old_embed.description, colour=KOALA_GREEN, channel=channel) msg_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) + rfr_msg_row = get_rfr_message(ctx.guild.id, channel.id, msg_id) else: await ctx.send("Okay, I'll stop the command then.") return @@ -487,7 +486,7 @@ 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.") @@ -542,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 @@ -561,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) @@ -571,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" @@ -657,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 @@ -674,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]) @@ -697,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 diff --git a/koala/cogs/react_for_role/db.py b/koala/cogs/react_for_role/db.py index 2314c754..ffb6351f 100644 --- a/koala/cogs/react_for_role/db.py +++ b/koala/cogs/react_for_role/db.py @@ -14,230 +14,228 @@ from .models import GuildRFRMessages, RFRMessageEmojiRoles, GuildRFRRequiredRoles from koala.db import assign_session -class ReactForRoleDBManager: - - @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)) +@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() - - @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 = ReactForRoleDBManager.get_rfr_message(guild_id, channel_id, message_id) - if not emoji_role_id: - return - else: - ReactForRoleDBManager.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() + 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] + +@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): + """ + 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] + +@assign_session +def get_rfr_reaction_role(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 - @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] - - @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]] = ReactForRoleDBManager.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): - """ - 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] - - @assign_session - def get_rfr_reaction_role(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 - - @assign_session - def get_rfr_reaction_role_by_emoji_str(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] - - @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 \ No newline at end of file +@assign_session +def get_rfr_reaction_role_by_emoji_str(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] + +@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 From 060667ec27ce98eab27dd9f90b606633766332a4 Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 31 Oct 2022 20:25:23 +0000 Subject: [PATCH 08/25] Fix db manager referencing in testing --- tests/cogs/react_for_role/test_cog.py | 2 +- tests/cogs/react_for_role/test_db.py | 1 - 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/tests/cogs/react_for_role/test_cog.py b/tests/cogs/react_for_role/test_cog.py index bdba553b..e89c76b1 100644 --- a/tests/cogs/react_for_role/test_cog.py +++ b/tests/cogs/react_for_role/test_cog.py @@ -590,7 +590,7 @@ async def test_rfr_remove_roles_from_msg(): 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() diff --git a/tests/cogs/react_for_role/test_db.py b/tests/cogs/react_for_role/test_db.py index 816ce90b..24f8503e 100644 --- a/tests/cogs/react_for_role/test_db.py +++ b/tests/cogs/react_for_role/test_db.py @@ -20,7 +20,6 @@ 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 3e932fe554de9619c541fbef51e18036c88caa2f Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 31 Oct 2022 20:32:10 +0000 Subject: [PATCH 09/25] Remove references to ReactForRole DBManager --- tests/cogs/react_for_role/test_db.py | 111 ++++++++++++++------------- tests/cogs/react_for_role/utils.py | 7 -- 2 files changed, 56 insertions(+), 62 deletions(-) diff --git a/tests/cogs/react_for_role/test_db.py b/tests/cogs/react_for_role/test_db.py index 24f8503e..94f4b7f8 100644 --- a/tests/cogs/react_for_role/test_db.py +++ b/tests/cogs/react_for_role/test_db.py @@ -24,7 +24,8 @@ 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 @@ -45,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) == [ @@ -56,12 +57,12 @@ 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, + 0] == get_rfr_message(guild2.id, channel2.id, msg_id) assert independent_get_guild_rfr_message(session) == expected_full_list @@ -69,23 +70,23 @@ async def test_rfr_db_functions_guild_rfr_messages(): 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) @@ -96,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) == [] @@ -108,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 @@ -116,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, + 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) @@ -236,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) @@ -265,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() @@ -275,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))): @@ -302,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): @@ -313,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..5426aeee 100644 --- a/tests/cogs/react_for_role/utils.py +++ b/tests/cogs/react_for_role/utils.py @@ -15,15 +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]]: sql_select = select(GuildRFRMessages) From b3b90d0896802f4d1602befdd5c16141d72f7a36 Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 31 Oct 2022 20:35:52 +0000 Subject: [PATCH 10/25] Remove DBManager references in test_cog.py --- tests/cogs/react_for_role/test_cog.py | 29 ++++++++++++++------------- 1 file changed, 15 insertions(+), 14 deletions(-) diff --git a/tests/cogs/react_for_role/test_cog.py b/tests/cogs/react_for_role/test_cog.py index e89c76b1..700bda0c 100644 --- a/tests/cogs/react_for_role/test_cog.py +++ b/tests/cogs/react_for_role/test_cog.py @@ -25,7 +25,8 @@ 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 @@ -58,7 +59,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)): @@ -343,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))): @@ -368,7 +369,7 @@ 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))): @@ -391,7 +392,7 @@ 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))): @@ -423,7 +424,7 @@ 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', @@ -448,7 +449,7 @@ 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', @@ -477,7 +478,7 @@ 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', @@ -504,8 +505,8 @@ 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"))] @@ -536,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 = [] @@ -570,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): @@ -580,7 +581,7 @@ 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', @@ -614,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 = [] From 8b7458e9a6898ba24d2b852ec4db3da453e0a010 Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 31 Oct 2022 20:50:55 +0000 Subject: [PATCH 11/25] Remove all import and remove a unused session import --- koala/cogs/react_for_role/cog.py | 2 +- koala/cogs/react_for_role/db.py | 1 - 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/koala/cogs/react_for_role/cog.py b/koala/cogs/react_for_role/cog.py index 01f404c5..a6da3822 100644 --- a/koala/cogs/react_for_role/cog.py +++ b/koala/cogs/react_for_role/cog.py @@ -23,7 +23,7 @@ from koala.colours import KOALA_GREEN from koala.utils import wait_for_message from koala.db import insert_extension -from .db import * +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 diff --git a/koala/cogs/react_for_role/db.py b/koala/cogs/react_for_role/db.py index ffb6351f..ac8d8f47 100644 --- a/koala/cogs/react_for_role/db.py +++ b/koala/cogs/react_for_role/db.py @@ -2,7 +2,6 @@ # Built-in/Generic Imports from typing import * -from requests import Session import sqlalchemy.exc import sqlalchemy.orm From 27904d0c78c92bedbbf19a413fd859641123bc2e Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 31 Oct 2022 21:00:16 +0000 Subject: [PATCH 12/25] Add session as a param to get_rfr_reaction_roles --- koala/cogs/react_for_role/db.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/koala/cogs/react_for_role/db.py b/koala/cogs/react_for_role/db.py index ac8d8f47..18215acd 100644 --- a/koala/cogs/react_for_role/db.py +++ b/koala/cogs/react_for_role/db.py @@ -171,7 +171,7 @@ def get_rfr_message_emoji_roles(emoji_role_id: int): 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): +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 From b439acb6af772347968379c419b1d4a7141a0276 Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 31 Oct 2022 21:03:51 +0000 Subject: [PATCH 13/25] Move the session param as added to wrong method --- koala/cogs/react_for_role/db.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/koala/cogs/react_for_role/db.py b/koala/cogs/react_for_role/db.py index 18215acd..bc6750df 100644 --- a/koala/cogs/react_for_role/db.py +++ b/koala/cogs/react_for_role/db.py @@ -158,7 +158,7 @@ def get_guild_rfr_roles(guild_id: int) -> List[int]: return role_ids @assign_session -def get_rfr_message_emoji_roles(emoji_role_id: int): +def get_rfr_message_emoji_roles(emoji_role_id: int, session: sqlalchemy.orm.Session): """ Returns all the emoji-role combinations on an rfr message @@ -171,7 +171,7 @@ def get_rfr_message_emoji_roles(emoji_role_id: int): 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): +def get_rfr_reaction_role(emoji_role_id: int, emoji_raw: str, role_id: int): """ Returns a specific emoji-role combo on an rfr message From 1f95586c9856afecf8f86f77b646fa094d208751 Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 31 Oct 2022 21:07:14 +0000 Subject: [PATCH 14/25] Add session as param to get_rfr_reaction_role --- koala/cogs/react_for_role/db.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/koala/cogs/react_for_role/db.py b/koala/cogs/react_for_role/db.py index bc6750df..aae2a8af 100644 --- a/koala/cogs/react_for_role/db.py +++ b/koala/cogs/react_for_role/db.py @@ -171,7 +171,7 @@ def get_rfr_message_emoji_roles(emoji_role_id: int, session: sqlalchemy.orm.Sess 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): +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 From 3c06626c1ce7ad44bf1e5cbc2cfbfcb434bf685f Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 31 Oct 2022 21:12:32 +0000 Subject: [PATCH 15/25] Add session as param to get_rfr_reaction_role_by_emoji_str --- koala/cogs/react_for_role/db.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/koala/cogs/react_for_role/db.py b/koala/cogs/react_for_role/db.py index aae2a8af..749e8ed7 100644 --- a/koala/cogs/react_for_role/db.py +++ b/koala/cogs/react_for_role/db.py @@ -189,7 +189,7 @@ def get_rfr_reaction_role(emoji_role_id: int, emoji_raw: str, role_id: int, sess return None @assign_session -def get_rfr_reaction_role_by_emoji_str(emoji_role_id: int, emoji_raw: str) -> Optional[int]: +def get_rfr_reaction_role_by_emoji_str(session: sqlalchemy.orm.Session, 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 From 28670301e606b4d7be4142f0b654681463be16b1 Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 31 Oct 2022 21:21:27 +0000 Subject: [PATCH 16/25] Move session to end of get_rfr_reaction_role_by_emoji_str --- koala/cogs/react_for_role/db.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/koala/cogs/react_for_role/db.py b/koala/cogs/react_for_role/db.py index 749e8ed7..69b28b1b 100644 --- a/koala/cogs/react_for_role/db.py +++ b/koala/cogs/react_for_role/db.py @@ -189,7 +189,7 @@ def get_rfr_reaction_role(emoji_role_id: int, emoji_raw: str, role_id: int, sess return None @assign_session -def get_rfr_reaction_role_by_emoji_str(session: sqlalchemy.orm.Session, emoji_role_id: int, emoji_raw: str) -> Optional[int]: +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 From 999ece385545a4ffd9ae07a8c5c513da902890de Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Tue, 1 Nov 2022 12:49:59 +0000 Subject: [PATCH 17/25] Write a test for test_get_first_emoji_from_str --- koala/cogs/react_for_role/core.py | 43 +++++++++++++++++++-------- koala/cogs/react_for_role/db.py | 42 ++++++++++++++++++-------- tests/cogs/react_for_role/test_cog.py | 23 ++++++++++++-- tests/cogs/react_for_role/test_db.py | 6 ++-- tests/cogs/react_for_role/utils.py | 4 ++- 5 files changed, 86 insertions(+), 32 deletions(-) diff --git a/koala/cogs/react_for_role/core.py b/koala/cogs/react_for_role/core.py index 6d6f8eac..a1c5b497 100644 --- a/koala/cogs/react_for_role/core.py +++ b/koala/cogs/react_for_role/core.py @@ -12,18 +12,22 @@ 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 } + return {'bot': bot, 'guild': guild} + @assign_session -async def create_rfr_message(title: str, guild: discord.Guild, description: str, colour: discord.Colour, channel: discord.TextChannel, **kwargs): +async def create_rfr_message(title: str, guild: discord.Guild, description: str, colour: discord.Colour, + channel: discord.TextChannel, **kwargs): embed: discord.Embed = discord.Embed(title=title, description=description, colour=colour) embed.set_footer(text="ReactForRole") embed.set_thumbnail(url=koala_logo) @@ -31,6 +35,7 @@ async def create_rfr_message(title: str, guild: discord.Guild, description: str, db.add_rfr_message(guild.id, channel.id, rfr_msg.id, **kwargs) return rfr_msg + @assign_session async def delete_rfr_message(guild_id: str, channel_id: str, msg: discord.Message, **kwargs): rfr_msg_row = db.get_rfr_message(guild_id, channel_id, msg.id, **kwargs) @@ -38,6 +43,7 @@ async def delete_rfr_message(guild_id: str, channel_id: str, msg: discord.Messag db.remove_rfr_message(guild_id, channel_id, msg.id, **kwargs) await msg.delete() + @assign_session async def use_inline_rfr_all(guild: discord.Guild, **kwargs): text_channels: List[discord.TextChannel] = guild.text_channels @@ -52,6 +58,7 @@ async def use_inline_rfr_all(guild: discord.Guild, **kwargs): embed.set_field_at(i, name=field.name, value=field.value, inline=True) await msg.edit(embed=embed) + async def use_inline_rfr_specific(embed: discord.Embed, msg: discord.Message): length = get_number_of_embed_fields(embed) for i in range(length): @@ -59,19 +66,23 @@ async def use_inline_rfr_specific(embed: discord.Embed, msg: discord.Message): embed.set_field_at(i, name=field.name, value=field.value, inline=True) await msg.edit(embed=embed) -async def rfr_edit(embed: discord.Embed, msg: discord.Message, description: str = "", title: str = "", image_url: str = ""): + +async def rfr_edit(embed: discord.Embed, msg: discord.Message, description: str = "", title: str = "", + image_url: str = ""): embed.description = description embed.title = title embed.set_thumbnail(url=image_url) await msg.edit(embed=embed) return msg + @assign_session -async def rfr_remove_emojis_roles(bot: Bot, guild: discord.Guild, msg: discord.Message, rfr_msg_row: discord.Message, wanted_removals: List[Union[discord.Emoji, str, discord.Role]], **kwargs): +async def rfr_remove_emojis_roles(bot: Bot, guild: discord.Guild, msg: discord.Message, rfr_msg_row: discord.Message, + 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 = discord.Embed(title=rfr_embed.title, description=rfr_embed.description, - colour=KOALA_GREEN) + colour=KOALA_GREEN) new_embed.set_thumbnail( url=koala_logo) new_embed.set_footer(text="ReactForRole") @@ -104,16 +115,18 @@ async def rfr_remove_emojis_roles(bot: Bot, guild: discord.Guild, msg: discord.M 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_add_emoji_role(guild: str, channel: discord.TextChannel, rfr_embed: discord.Embed, msg: discord.Message, rfr_msg_row: discord.Message, emoji_role_map: List[Tuple[Union[discord.Emoji, str], discord.Role]], **kwargs): +async def rfr_add_emoji_role(guild: str, channel: discord.TextChannel, rfr_embed: discord.Embed, msg: discord.Message, + rfr_msg_row: discord.Message, + emoji_role_map: List[Tuple[Union[discord.Emoji, str], discord.Role]], **kwargs): duplicateRolesFound = False duplicateEmojisFound = False @@ -122,13 +135,13 @@ async def rfr_add_emoji_role(guild: str, channel: discord.TextChannel, rfr_embed role = emoji_role[1] if discord_emoji in [x.name for x in rfr_embed.fields]: - duplicateEmojisFound = True + duplicateEmojisFound = True elif role in [x.value for x in rfr_embed.fields]: - duplicateRolesFound = True + duplicateRolesFound = True else: if isinstance(discord_emoji, str): db.add_rfr_message_emoji_role(rfr_msg_row[3], emoji.demojize(discord_emoji), - role.id, **kwargs) + 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) @@ -146,21 +159,25 @@ async def rfr_add_emoji_role(guild: str, channel: discord.TextChannel, rfr_embed edited_msg = await msg.edit(embed=rfr_embed) return duplicateRolesFound, duplicateEmojisFound, edited_msg + async def add_guild_rfr_required_role(bot: Bot, guild: discord.Guild, role_str: str, **kwargs): ctx = create_ctx(bot, guild) role: discord.Role = await commands.RoleConverter().convert(ctx, role_str) db.remove_guild_rfr_required_role(ctx.guild.id, role.id, **kwargs) return role + async def remove_guild_rfr_required_role(bot: Bot, guild: discord.Guild, role_str: str, **kwargs): ctx = create_ctx(bot, guild) role: discord.Role = await commands.RoleConverter().convert(ctx, role_str) db.add_guild_rfr_required_role(ctx.guild.id, role.id, **kwargs) return role + def rfr_list_guild_required_roles(guild: discord.Guild, **kwargs): return 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 @@ -181,6 +198,7 @@ async def setup_rfr_reaction_permissions(guild: discord.Guild, channel: discord. 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 @@ -197,6 +215,7 @@ def get_embed_from_message(msg: discord.Message) -> Optional[discord.Embed]: except IndexError: return None + def get_number_of_embed_fields(embed: discord.Embed) -> int: """ Gets the number of fields in an embed. @@ -234,4 +253,4 @@ async def get_first_emoji_from_str(bot: Bot, guild: discord.Guild, content: str) 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?" \ No newline at end of file + 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 69b28b1b..63617bd1 100644 --- a/koala/cogs/react_for_role/db.py +++ b/koala/cogs/react_for_role/db.py @@ -13,6 +13,7 @@ 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): """ @@ -27,6 +28,7 @@ def add_rfr_message(guild_id: int, channel_id: int, message_id: int, session: sq 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): """ @@ -41,10 +43,12 @@ def add_rfr_message_emoji_role(emoji_role_id: int, emoji_raw: str, role_id: int, session.commit() except sqlalchemy.exc.IntegrityError: logger.warning("RFRMessageEmojiRoles already exists for <%s, %s, %s>, continuing", - emoji_role_id, emoji_raw, role_id) + 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): +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 @@ -54,14 +58,14 @@ def remove_rfr_message_emoji_role(emoji_role_id: int, emoji_raw: str = None, rol :return: """ if not emoji_raw: - delete_sql = delete(RFRMessageEmojiRoles)\ + delete_sql = delete(RFRMessageEmojiRoles) \ .where( and_( RFRMessageEmojiRoles.emoji_role_id == emoji_role_id, RFRMessageEmojiRoles.role_id == role_id )) else: - delete_sql = delete(RFRMessageEmojiRoles)\ + delete_sql = delete(RFRMessageEmojiRoles) \ .where( and_( RFRMessageEmojiRoles.emoji_role_id == emoji_role_id, @@ -70,6 +74,7 @@ def remove_rfr_message_emoji_role(emoji_role_id: int, emoji_raw: str = None, rol session.execute(delete_sql) session.commit() + @assign_session def remove_rfr_message_emoji_roles(emoji_role_id: int, session: sqlalchemy.orm.Session): """ @@ -83,6 +88,7 @@ def remove_rfr_message_emoji_roles(emoji_role_id: int, session: sqlalchemy.orm.S 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): """ @@ -100,14 +106,16 @@ def remove_rfr_message(guild_id: int, channel_id: int, message_id: int, session: delete_sql = delete(GuildRFRMessages) \ .where(and_(and_( - GuildRFRMessages.guild_id == guild_id, - GuildRFRMessages.channel_id == channel_id), - GuildRFRMessages.message_id == message_id)) + 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]]: +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 @@ -116,14 +124,15 @@ def get_rfr_message(guild_id: int, channel_id: int, message_id: int, session: sq :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() + .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]]: """ @@ -132,10 +141,11 @@ def get_guild_rfr_messages(guild_id: int, session: sqlalchemy.orm.Session) -> Li :return: List of rfr messages in the guild. """ messages = session.execute(select(GuildRFRMessages) - .filter_by(guild_id=guild_id)).scalars().all() + .filter_by(guild_id=guild_id)).scalars().all() return [message.old_format() for message in messages] + @assign_session def get_guild_rfr_roles(guild_id: int) -> List[int]: """ @@ -157,6 +167,7 @@ def get_guild_rfr_roles(guild_id: int) -> List[int]: role_ids.extend(ids) return role_ids + @assign_session def get_rfr_message_emoji_roles(emoji_role_id: int, session: sqlalchemy.orm.Session): """ @@ -170,6 +181,7 @@ def get_rfr_message_emoji_roles(emoji_role_id: int, session: sqlalchemy.orm.Sess 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): """ @@ -188,6 +200,7 @@ def get_rfr_reaction_role(emoji_role_id: int, emoji_raw: str, role_id: int, sess else: return None + @assign_session def get_rfr_reaction_role_by_emoji_str(emoji_role_id: int, emoji_raw: str, session: sqlalchemy.orm.Session): """ @@ -198,11 +211,12 @@ def get_rfr_reaction_role_by_emoji_str(emoji_role_id: int, emoji_raw: str, sessi """ 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() + .filter_by(emoji_role_id=emoji_role_id, emoji_raw=emoji_raw)).one_or_none() if not row: return return row[0] + @assign_session def add_guild_rfr_required_role(guild_id: int, role_id: int, session: sqlalchemy.orm.Session): """ @@ -214,6 +228,7 @@ def add_guild_rfr_required_role(guild_id: int, role_id: int, session: sqlalchemy 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): """ @@ -225,6 +240,7 @@ def remove_guild_rfr_required_role(guild_id: int, role_id: int, session: sqlalch 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]: """ diff --git a/tests/cogs/react_for_role/test_cog.py b/tests/cogs/react_for_role/test_cog.py index 700bda0c..dede35fa 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 @@ -30,6 +28,7 @@ from tests.log import logger from koala.cogs import ReactForRole + # Constants # Variables @@ -68,7 +67,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): @@ -630,3 +628,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 94f4b7f8..77fd5b06 100644 --- a/tests/cogs/react_for_role/test_db.py +++ b/tests/cogs/react_for_role/test_db.py @@ -63,8 +63,8 @@ async def test_rfr_db_functions_guild_rfr_messages(): expected_full_list[1]] assert independent_get_guild_rfr_message(session, guild2.id, channel2.id, msg_id)[ 0] == get_rfr_message(guild2.id, - channel2.id, - msg_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( @@ -158,7 +158,7 @@ async def test_rfr_db_functions_rfr_message_emoji_roles(): 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] == get_rfr_reaction_role_by_emoji_str(1, - fake_emoji_1) + fake_emoji_1) assert independent_get_rfr_message_emoji_role(session, 1) == get_rfr_message_emoji_roles(1) assert independent_get_rfr_message_emoji_role(session, 1, role_id=fake_role_id_2)[0][ diff --git a/tests/cogs/react_for_role/utils.py b/tests/cogs/react_for_role/utils.py index 5426aeee..b3a0d259 100644 --- a/tests/cogs/react_for_role/utils.py +++ b/tests/cogs/react_for_role/utils.py @@ -17,6 +17,7 @@ # Own modules from koala.cogs.react_for_role.models import GuildRFRRequiredRoles, GuildRFRMessages, RFRMessageEmojiRoles + 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]]: sql_select = select(GuildRFRMessages) @@ -30,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: From 09e2aa5e538d7caf4d976d4df92b96f5a73c26b9 Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Sun, 20 Nov 2022 17:54:28 +0000 Subject: [PATCH 18/25] Try statement for setting permissions --- koala/cogs/react_for_role/cog.py | 3 +-- koala/cogs/react_for_role/core.py | 7 +++++-- tests/cogs/react_for_role/test_cog.py | 3 +-- 3 files changed, 7 insertions(+), 6 deletions(-) diff --git a/koala/cogs/react_for_role/cog.py b/koala/cogs/react_for_role/cog.py index a6da3822..ff301c2d 100644 --- a/koala/cogs/react_for_role/cog.py +++ b/koala/cogs/react_for_role/cog.py @@ -182,8 +182,7 @@ async def rfr_create_message(self, ctx: commands.Context): f"I'll create the react for role message now.") rfr_msg = await core.create_rfr_message(title, ctx.guild, desc, KOALA_GREEN, channel) - # TODO - Get this working, for some reason we get 403 currently - # await core.setup_rfr_reaction_permissions(ctx.guild, channel, self.bot) + 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 " diff --git a/koala/cogs/react_for_role/core.py b/koala/cogs/react_for_role/core.py index a1c5b497..cb5b8a8b 100644 --- a/koala/cogs/react_for_role/core.py +++ b/koala/cogs/react_for_role/core.py @@ -191,8 +191,11 @@ async def setup_rfr_reaction_permissions(guild: discord.Guild, channel: discord. 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) + try: + await channel.set_permissions(role, overwrite=overwrite) + except discord.Forbidden: + logger.error(f"ReactForRole: Failed to set permissions for channel {channel.id} in guild {guild.id}.") + return False 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: diff --git a/tests/cogs/react_for_role/test_cog.py b/tests/cogs/react_for_role/test_cog.py index dede35fa..df78a669 100644 --- a/tests/cogs/react_for_role/test_cog.py +++ b/tests/cogs/react_for_role/test_cog.py @@ -641,9 +641,8 @@ async def test_get_first_emoji_from_str(): 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) + assert core.get_first_emoji_from_str(koalabot, guild, uni_emoji) \ No newline at end of file From b46a23b644443598bf80d0a4aeacde7a47f9065e Mon Sep 17 00:00:00 2001 From: Otto Hooper Date: Mon, 13 Mar 2023 19:05:04 +0000 Subject: [PATCH 19/25] Revert "Try statement for setting permissions" This reverts commit 09e2aa5e538d7caf4d976d4df92b96f5a73c26b9. --- koala/cogs/react_for_role/cog.py | 3 ++- koala/cogs/react_for_role/core.py | 7 ++----- tests/cogs/react_for_role/test_cog.py | 3 ++- 3 files changed, 6 insertions(+), 7 deletions(-) diff --git a/koala/cogs/react_for_role/cog.py b/koala/cogs/react_for_role/cog.py index ff301c2d..a6da3822 100644 --- a/koala/cogs/react_for_role/cog.py +++ b/koala/cogs/react_for_role/cog.py @@ -182,7 +182,8 @@ async def rfr_create_message(self, ctx: commands.Context): f"I'll create the react for role message now.") rfr_msg = await core.create_rfr_message(title, ctx.guild, desc, KOALA_GREEN, channel) - await core.setup_rfr_reaction_permissions(ctx.guild, channel, self.bot) + # 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 " diff --git a/koala/cogs/react_for_role/core.py b/koala/cogs/react_for_role/core.py index cb5b8a8b..a1c5b497 100644 --- a/koala/cogs/react_for_role/core.py +++ b/koala/cogs/react_for_role/core.py @@ -191,11 +191,8 @@ async def setup_rfr_reaction_permissions(guild: discord.Guild, channel: discord. role: discord.Role = discord.utils.get(guild.roles, id=guild.id) overwrite: discord.PermissionOverwrite = discord.PermissionOverwrite() overwrite.update(add_reactions=False) - try: - await channel.set_permissions(role, overwrite=overwrite) - except discord.Forbidden: - logger.error(f"ReactForRole: Failed to set permissions for channel {channel.id} in guild {guild.id}.") - return 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: diff --git a/tests/cogs/react_for_role/test_cog.py b/tests/cogs/react_for_role/test_cog.py index df78a669..dede35fa 100644 --- a/tests/cogs/react_for_role/test_cog.py +++ b/tests/cogs/react_for_role/test_cog.py @@ -641,8 +641,9 @@ async def test_get_first_emoji_from_str(): 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) \ No newline at end of file + assert core.get_first_emoji_from_str(koalabot, guild, uni_emoji) From 8504925f25f8437be65bc5300cad20c87adaff1b Mon Sep 17 00:00:00 2001 From: JayDwee Date: Sat, 18 Mar 2023 19:00:33 +0000 Subject: [PATCH 20/25] chore: add rfr create API --- koala/cogs/react_for_role/__init__.py | 10 ++-- koala/cogs/react_for_role/api.py | 67 +++++++++++++++++++++++++++ koala/cogs/react_for_role/cog.py | 11 ++--- koala/cogs/react_for_role/core.py | 2 +- koala/rest/api.py | 14 ++++-- tests/cogs/react_for_role/test_api.py | 35 ++++++++++++++ tests/cogs/react_for_role/test_cog.py | 22 +++++++++ 7 files changed, 146 insertions(+), 15 deletions(-) create mode 100644 koala/cogs/react_for_role/api.py create mode 100644 tests/cogs/react_for_role/test_api.py diff --git a/koala/cogs/react_for_role/__init__.py b/koala/cogs/react_for_role/__init__.py index 62419b41..e8443764 100644 --- a/koala/cogs/react_for_role/__init__.py +++ b/koala/cogs/react_for_role/__init__.py @@ -1,5 +1,7 @@ -from . import utils, db, models, cog, core -from .cog import ReactForRole, setup +from . import utils, db, models, cog, core, api +from .cog import ReactForRole -def setup(bot): - cog.setup(bot) \ No newline at end of file + +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..0e587b50 --- /dev/null +++ b/koala/cogs/react_for_role/api.py @@ -0,0 +1,67 @@ +# Futures +# Built-in/Generic Imports +# Libs +from http.client import CREATED, OK, BAD_REQUEST +from aiohttp import web +import discord +from discord.ext import commands +from discord.ext.commands import Bot + +# Own modules +from . import core +from .log import logger +from koala.rest.api import parse_request, build_response +from koala.utils import convert_iso_datetime + +# Constants +RFR_ENDPOINT = 'rfr' +CREATE = 'create' + + +class RfrEndpoint: + _bot: commands.Bot + """ + The API endpoints for BaseCog + """ + def __init__(self, bot): + self._bot = bot + + def register(self, app): + """ + Register the routes for the given application + todo: review aiohttp 'views' and see if they are a better idea + :param app: The aiohttp.web.Application (likely of the sub app) + :return: app + """ + app.add_routes([web.post('/{}'.format(CREATE), self.post_create_rfr_message)]) + return app + + @parse_request + async def post_create_rfr_message(self, guild_id: int, channel_id: int, + title: str, description: str, colour: str): + """ + Create a React For Role message + + :param guild_id: ID of guild + :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 + :return: + """ + return {"rfr_id": await core.create_rfr_message(title=title, guild=self._bot.get_guild(guild_id), + description=description, + colour=discord.Colour.from_str(colour), + channel=self._bot.get_channel(channel_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 5acf9d2d..0804740c 100644 --- a/koala/cogs/react_for_role/cog.py +++ b/koala/cogs/react_for_role/cog.py @@ -181,12 +181,12 @@ async def rfr_create_message(self, ctx: commands.Context): 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.") - rfr_msg = await core.create_rfr_message(title, ctx.guild, desc, KOALA_GREEN, channel) + rfr_msg_id = await core.create_rfr_message(title, ctx.guild, desc, KOALA_GREEN, channel) # 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() @@ -444,10 +444,9 @@ async def rfr_add_roles_to_msg(self, ctx: commands.Context): "Okay, I'll continue then. The new message will have the same title and description as the " "old one.") old_embed = core.get_embed_from_message(msg) - msg = core.create_rfr_message(title=old_embed.title, guild=ctx.guild, description=old_embed.description, colour=KOALA_GREEN, channel=channel) - msg_id = msg.id - await ctx.send(f"Okay, the new message has ID {msg.id} and is in {msg.channel.mention}.") - rfr_msg_row = get_rfr_message(ctx.guild.id, channel.id, msg_id) + rfr_msg_id = core.create_rfr_message(title=old_embed.title, guild=ctx.guild, description=old_embed.description, colour=KOALA_GREEN, channel=channel) + await ctx.send(f"Okay, the new message has ID {rfr_msg_id} and is in {msg.channel.mention}.") + rfr_msg_row = get_rfr_message(ctx.guild.id, channel.id, rfr_msg_id) else: await ctx.send("Okay, I'll stop the command then.") return diff --git a/koala/cogs/react_for_role/core.py b/koala/cogs/react_for_role/core.py index a1c5b497..d054dc4b 100644 --- a/koala/cogs/react_for_role/core.py +++ b/koala/cogs/react_for_role/core.py @@ -33,7 +33,7 @@ async def create_rfr_message(title: str, guild: discord.Guild, description: str, embed.set_thumbnail(url=koala_logo) rfr_msg: discord.Message = await channel.send(embed=embed) db.add_rfr_message(guild.id, channel.id, rfr_msg.id, **kwargs) - return rfr_msg + return rfr_msg.id @assign_session diff --git a/koala/rest/api.py b/koala/rest/api.py index 76099167..aec548af 100644 --- a/koala/rest/api.py +++ b/koala/rest/api.py @@ -8,6 +8,7 @@ # Libs from functools import wraps import aiohttp.web +from aiohttp.abc import Request # Own modules from koala.models import BaseModel @@ -40,8 +41,13 @@ def build_response(status_code, data): :param data: :return: """ + if data: + body = json.dumps(data, cls=EnhancedJSONEncoder) + else: + body = None + return aiohttp.web.Response(status=status_code, - body=json.dumps(data, cls=EnhancedJSONEncoder), + body=body, content_type='application/json') @@ -84,15 +90,15 @@ def parsed_request(func): @wraps(func) async def wrapper(*args, **kwargs): self = args[0] - request = args[1] + request: Request = args[1] wanted_args = list(inspect.signature(func).parameters.keys()) wanted_args.remove("self") available_args = {} - if (request.method == "POST" or request.method == "PUT") and request.has_body: - body = await request.post() + if (request.method == "POST" or request.method == "PUT") and request.can_read_body: + body = await request.json() for arg in wanted_args: if arg in body: available_args[arg] = body[arg] 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..5ee5ebf3 --- /dev/null +++ b/tests/cogs/react_for_role/test_api.py @@ -0,0 +1,35 @@ +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 + +# 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_create_rfr(api_client): + resp = await api_client.post('/create', 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 "rfr_id" in resp_json.keys() diff --git a/tests/cogs/react_for_role/test_cog.py b/tests/cogs/react_for_role/test_cog.py index 4a72d127..ee861c3f 100644 --- a/tests/cogs/react_for_role/test_cog.py +++ b/tests/cogs/react_for_role/test_cog.py @@ -383,6 +383,28 @@ async def test_rfr_edit_description(): assert dpytest.verify().message() +@pytest.mark.asyncio +async def test_rfr_edit_inline(): + config: dpytest.RunnerConfig = dpytest.get_config() + guild: discord.Guild = config.guilds[0] + channel: discord.TextChannel = guild.text_channels[0] + embed: discord.Embed = discord.Embed(title="title", description="description") + client: discord.Client = config.client + message: discord.Message = await dpytest.message("rfr") + msg_id = message.id + add_rfr_message(guild.id, channel.id, msg_id) + assert embed.fields[0].inline == True + 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.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() + assert dpytest.verify().message() + assert dpytest.verify().message() + @pytest.mark.asyncio async def test_rfr_edit_title(): config: dpytest.RunnerConfig = dpytest.get_config() From 01f39cd58c21ee79caf205ca968cf52c02e25db6 Mon Sep 17 00:00:00 2001 From: JayDwee Date: Tue, 21 Mar 2023 19:38:36 +0000 Subject: [PATCH 21/25] fix: breaking tests --- koala/rest/api.py | 2 +- tests/cogs/base/test_api.py | 108 +++++++++++++------------- tests/cogs/react_for_role/test_cog.py | 22 ------ 3 files changed, 55 insertions(+), 77 deletions(-) diff --git a/koala/rest/api.py b/koala/rest/api.py index aec548af..b44c705e 100644 --- a/koala/rest/api.py +++ b/koala/rest/api.py @@ -41,7 +41,7 @@ def build_response(status_code, data): :param data: :return: """ - if data: + if data is not None: body = json.dumps(data, cls=EnhancedJSONEncoder) else: body = None diff --git a/tests/cogs/base/test_api.py b/tests/cogs/base/test_api.py index 34addabb..2eb6ef7b 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', data=( + 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', data=( + 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', data=( + 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', data=( + 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', data=( + 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', data=( + 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', data=( + 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', data=( + 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,54 @@ async def test_get_support_link(api_client): ''' async def test_post_load_cog(api_client): - resp = await api_client.post('/load-cog', data=( + resp = await api_client.post('/load-cog', json= { 'extension': 'announce', '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', data=( + resp = await api_client.post('/load-cog', json= { 'extension': 'base', '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', data=( + resp = await api_client.post('/load-cog', json= { 'extension': 'invalidCog', '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', data=( + 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_already_loaded(api_client): - await api_client.post('/load-cog', data=( + await api_client.post('/load-cog', json= { 'extension': 'announce', 'package': koalabot.COGS_PACKAGE - })) + }) - resp = await api_client.post('/load-cog', data=( + resp = await api_client.post('/load-cog', json= { 'extension': 'announce', 'package': koalabot.COGS_PACKAGE - })) + }) assert resp.status == UNPROCESSABLE_ENTITY assert await resp.text() == '422: Error loading cog: Already loaded' @@ -261,44 +261,44 @@ async def test_post_load_cog_already_loaded(api_client): ''' async def test_post_unload_cog(api_client): - await api_client.post('/load-cog', data=( + await api_client.post('/load-cog', json= { 'extension': 'announce', 'package': koalabot.COGS_PACKAGE - })) + }) - resp = await api_client.post('/unload-cog', data=( + resp = await api_client.post('/unload-cog', json= { 'extension': 'announce', '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', data=( + resp = await api_client.post('/unload-cog', json= { 'extension': 'announce', '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', data=( + 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_base_cog(api_client): - resp = await api_client.post('/unload-cog', data=( + resp = await api_client.post('/unload-cog', json= { 'extension': 'BaseCog', '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 +313,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', data=({ + 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 +325,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', data=( + 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', data=( + 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 +354,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', data=({ + 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', data=({ + 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', data=({ + 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', data=({ + 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 +390,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', data=( + 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_cog.py b/tests/cogs/react_for_role/test_cog.py index ee861c3f..4a72d127 100644 --- a/tests/cogs/react_for_role/test_cog.py +++ b/tests/cogs/react_for_role/test_cog.py @@ -383,28 +383,6 @@ async def test_rfr_edit_description(): assert dpytest.verify().message() -@pytest.mark.asyncio -async def test_rfr_edit_inline(): - config: dpytest.RunnerConfig = dpytest.get_config() - guild: discord.Guild = config.guilds[0] - channel: discord.TextChannel = guild.text_channels[0] - embed: discord.Embed = discord.Embed(title="title", description="description") - client: discord.Client = config.client - message: discord.Message = await dpytest.message("rfr") - msg_id = message.id - add_rfr_message(guild.id, channel.id, msg_id) - assert embed.fields[0].inline == True - 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.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() - assert dpytest.verify().message() - assert dpytest.verify().message() - @pytest.mark.asyncio async def test_rfr_edit_title(): config: dpytest.RunnerConfig = dpytest.get_config() From e8c5808e86c279b7dfc4c8e30e66dcee43de39be Mon Sep 17 00:00:00 2001 From: JayDwee Date: Tue, 28 Mar 2023 21:48:45 +0100 Subject: [PATCH 22/25] chore: add all RFR endpoints and tests --- koala/cogs/react_for_role/api.py | 207 +++++++++++++++++++++-- koala/cogs/react_for_role/cog.py | 44 ++--- koala/cogs/react_for_role/core.py | 231 +++++++++++++++++++------- koala/cogs/react_for_role/db.py | 24 +-- koala/cogs/react_for_role/dto.py | 32 ++++ koala/rest/api.py | 34 ++-- tests/cogs/react_for_role/test_api.py | 223 +++++++++++++++++++++++-- 7 files changed, 668 insertions(+), 127 deletions(-) create mode 100644 koala/cogs/react_for_role/dto.py diff --git a/koala/cogs/react_for_role/api.py b/koala/cogs/react_for_role/api.py index 0e587b50..d915d71b 100644 --- a/koala/cogs/react_for_role/api.py +++ b/koala/cogs/react_for_role/api.py @@ -1,58 +1,231 @@ # Futures # Built-in/Generic Imports # Libs -from http.client import CREATED, OK, BAD_REQUEST -from aiohttp import web +from typing import List + import discord -from discord.ext import commands +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 koala.rest.api import parse_request, build_response -from koala.utils import convert_iso_datetime +from ... import colours # Constants -RFR_ENDPOINT = 'rfr' -CREATE = 'create' +RFR_ENDPOINT = 'react-for-role' + +MESSAGE = 'message' +REQUIRED_ROLES = 'required-roles' class RfrEndpoint: - _bot: commands.Bot + _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 - todo: review aiohttp 'views' and see if they are a better idea :param app: The aiohttp.web.Application (likely of the sub app) :return: app """ - app.add_routes([web.post('/{}'.format(CREATE), self.post_create_rfr_message)]) + 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_create_rfr_message(self, guild_id: int, channel_id: int, - title: str, description: str, colour: str): + 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: ID of guild + :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 {"rfr_id": await core.create_rfr_message(title=title, guild=self._bot.get_guild(guild_id), - description=description, - colour=discord.Colour.from_str(colour), - channel=self._bot.get_channel(channel_id))} + return core.rfr_list_guild_required_roles(self._bot.get_guild(int(guild_id))) def setup(bot: Bot): diff --git a/koala/cogs/react_for_role/cog.py b/koala/cogs/react_for_role/cog.py index 0804740c..ea8340f9 100644 --- a/koala/cogs/react_for_role/cog.py +++ b/koala/cogs/react_for_role/cog.py @@ -181,7 +181,7 @@ async def rfr_create_message(self, ctx: commands.Context): 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.") - rfr_msg_id = await core.create_rfr_message(title, ctx.guild, desc, KOALA_GREEN, channel) + 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) @@ -206,7 +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") - await core.delete_rfr_message(ctx.guild.id, channel.id, msg) + 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,7 +234,7 @@ async def rfr_edit_description(self, ctx: commands.Context): 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": - await core.rfr_edit(embed, msg, description=desc) + await core.rfr_edit(msg, description=desc) else: await ctx.send("Okay, cancelling command.") else: @@ -259,7 +259,7 @@ async def rfr_edit_title(self, ctx: commands.Context): 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": - await core.rfr_edit(embed, msg, title=title) + await core.rfr_edit(msg, title=title) else: await ctx.send("Okay, cancelling command.") else: @@ -292,13 +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: - await core.rfr_edit(embed, msg, image_url=str(image.url)) + 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) - await core.rfr_edit(embed, msg, image_url=img_url) + 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.") @@ -361,7 +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.") - await core.use_inline_rfr_specific(embed, msg) + await core.use_inline_rfr_specific(msg) await ctx.send("Okay, should be done. Please check.") @commands.check(koalabot.is_admin) @@ -444,9 +444,10 @@ async def rfr_add_roles_to_msg(self, ctx: commands.Context): "Okay, I'll continue then. The new message will have the same title and description as the " "old one.") old_embed = core.get_embed_from_message(msg) - rfr_msg_id = core.create_rfr_message(title=old_embed.title, guild=ctx.guild, description=old_embed.description, colour=KOALA_GREEN, channel=channel) + 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}.") - rfr_msg_row = get_rfr_message(ctx.guild.id, channel.id, rfr_msg_id) else: await ctx.send("Okay, I'll stop the command then.") return @@ -460,8 +461,8 @@ 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 = core.get_embed_from_message(msg) - duplicateRolesFound, duplicateEmojisFound, edited_msg = core.rfr_add_emoji_role(ctx.guild, channel, rfr_embed, msg, rfr_msg_row, emoji_role_list) + 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.") @@ -499,7 +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.") - await core.delete_rfr_message(ctx.guild.id, channel.id, msg) + 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: @@ -524,7 +525,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, I'll delete the message then.") - await core.delete_rfr_message(ctx.guild.id, channel.id, msg) + await core.delete_rfr_message(self.bot, msg.id, ctx.guild.id, channel.id) return await ctx.send("Okay, I've removed those options from the react for role message.") @@ -586,17 +587,17 @@ 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 core.add_guild_rfr_required_role(self.bot, ctx.guild, 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.") except (commands.CommandError, commands.BadArgument): await ctx.send("Found an issue with your provided argument, couldn't get an actual role. Please try again.") @@ -604,17 +605,17 @@ async def rfr_add_guild_required_role(self, ctx: commands.Context, role_str: str @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 core.remove_guild_rfr_required_role(self.bot, ctx.guild, 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.") except (commands.CommandError, commands.BadArgument): @@ -630,7 +631,7 @@ async def rfr_list_guild_required_roles(self, ctx: commands.Context): :param ctx: Context of the command. :return: """ - role_ids = core.rfr_list_guild_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: @@ -745,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 diff --git a/koala/cogs/react_for_role/core.py b/koala/cogs/react_for_role/core.py index d054dc4b..fcd58dc6 100644 --- a/koala/cogs/react_for_role/core.py +++ b/koala/cogs/react_for_role/core.py @@ -6,7 +6,10 @@ 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 @@ -26,22 +29,88 @@ def create_ctx(bot: Bot, guild: discord.Guild): @assign_session -async def create_rfr_message(title: str, guild: discord.Guild, description: str, colour: discord.Colour, - channel: discord.TextChannel, **kwargs): +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") - embed.set_thumbnail(url=koala_logo) + 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) - return rfr_msg.id + 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 delete_rfr_message(guild_id: str, channel_id: str, msg: discord.Message, **kwargs): - rfr_msg_row = db.get_rfr_message(guild_id, channel_id, msg.id, **kwargs) +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, msg.id, **kwargs) - await msg.delete() + 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 @@ -50,7 +119,7 @@ async def use_inline_rfr_all(guild: discord.Guild, **kwargs): 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(id=rfr_message[2]) + 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): @@ -59,33 +128,37 @@ async def use_inline_rfr_all(guild: discord.Guild, **kwargs): await msg.edit(embed=embed) -async def use_inline_rfr_specific(embed: discord.Embed, msg: discord.Message): - length = get_number_of_embed_fields(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 = embed.fields[i] - embed.set_field_at(i, name=field.name, value=field.value, inline=True) - await msg.edit(embed=embed) + 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(embed: discord.Embed, msg: discord.Message, description: str = "", title: str = "", - image_url: str = ""): - embed.description = description - embed.title = title - embed.set_thumbnail(url=image_url) - await msg.edit(embed=embed) - return msg +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: discord.Message, +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 = discord.Embed(title=rfr_embed.title, description=rfr_embed.description, - colour=KOALA_GREEN) - new_embed.set_thumbnail( - url=koala_logo) - new_embed.set_footer(text="ReactForRole") + new_embed = rfr_embed.copy() + new_embed.clear_fields() removed_field_indexes = [] reactions_to_remove: List[discord.Reaction] = [] errors = [] @@ -104,7 +177,7 @@ async def rfr_remove_emojis_roles(bot: Bot, guild: discord.Guild, msg: discord.M field = rfr_embed_fields[field_index] removed_field_indexes.append(field_index) - reaction_emoji, err = get_first_emoji_from_str(bot, guild, field.name) + 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] @@ -124,20 +197,51 @@ async def rfr_remove_emojis_roles(bot: Bot, guild: discord.Guild, msg: discord.M @assign_session -async def rfr_add_emoji_role(guild: str, channel: discord.TextChannel, rfr_embed: discord.Embed, msg: discord.Message, - rfr_msg_row: discord.Message, - emoji_role_map: List[Tuple[Union[discord.Emoji, str], discord.Role]], **kwargs): - duplicateRolesFound = False - duplicateEmojisFound = False +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]: - duplicateEmojisFound = True + duplicate_emojis_found = True elif role in [x.value for x in rfr_embed.fields]: - duplicateRolesFound = True + duplicate_roles_found = True else: if isinstance(discord_emoji, str): db.add_rfr_message_emoji_role(rfr_msg_row[3], emoji.demojize(discord_emoji), @@ -150,32 +254,46 @@ async def rfr_add_emoji_role(guild: str, channel: discord.TextChannel, rfr_embed 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 {discord_emoji}.") + 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 duplicateRolesFound, duplicateEmojisFound, edited_msg + return duplicate_roles_found, duplicate_emojis_found, edited_msg -async def add_guild_rfr_required_role(bot: Bot, guild: discord.Guild, role_str: str, **kwargs): - ctx = create_ctx(bot, guild) - role: discord.Role = await commands.RoleConverter().convert(ctx, role_str) - db.remove_guild_rfr_required_role(ctx.guild.id, role.id, **kwargs) - return role +@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) -async def remove_guild_rfr_required_role(bot: Bot, guild: discord.Guild, role_str: str, **kwargs): - ctx = create_ctx(bot, guild) - role: discord.Role = await commands.RoleConverter().convert(ctx, role_str) - db.add_guild_rfr_required_role(ctx.guild.id, role.id, **kwargs) - return role + +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 db.get_guild_rfr_required_roles(guild.id, **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): @@ -225,18 +343,17 @@ def get_number_of_embed_fields(embed: discord.Embed) -> int: return len(embed.fields) -async def get_first_emoji_from_str(bot: Bot, guild: discord.Guild, content: str) -> Optional[ - Union[discord.Emoji, str]]: +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 ctx: Context of the original command + :param bot: + :param guild: :param content: Message content :return: Emoji if there is a valid one. Otherwise None. """ - ctx = create_ctx(bot, guild) - # First check for a custom discord emoji in the string search_result = CUSTOM_EMOJI_REGEXP.search(str(content)) if not search_result: @@ -246,9 +363,11 @@ async def get_first_emoji_from_str(bot: Bot, guild: discord.Guild, content: str) return None, "No emoji found." return content, None else: - emoji_str = search_result.group().strip() + emoji_id = int(search_result[:-1].split(":")[-1]) try: - discord_emoji: discord.Emoji = await commands.EmojiConverter().convert(ctx, emoji_str) + 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." diff --git a/koala/cogs/react_for_role/db.py b/koala/cogs/react_for_role/db.py index 63617bd1..46c6aa0a 100644 --- a/koala/cogs/react_for_role/db.py +++ b/koala/cogs/react_for_role/db.py @@ -169,36 +169,36 @@ def get_guild_rfr_roles(guild_id: int) -> List[int]: @assign_session -def get_rfr_message_emoji_roles(emoji_role_id: int, session: sqlalchemy.orm.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 """ - with session_manager() as session: - rows = session.execute(select(RFRMessageEmojiRoles).filter_by(emoji_role_id=emoji_role_id)).scalars().all() + 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] + 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): +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 """ - 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 + 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 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 b44c705e..e718bef9 100644 --- a/koala/rest/api.py +++ b/koala/rest/api.py @@ -7,9 +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 @@ -17,6 +21,7 @@ from http.client import OK + # Variables @@ -24,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() @@ -47,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 @@ -92,31 +98,39 @@ async def wrapper(*args, **kwargs): self = args[0] request: Request = args[1] - wanted_args = list(inspect.signature(func).parameters.keys()) - wanted_args.remove("self") + wanted_args: dict[str, inspect.Parameter] = dict(inspect.signature(func).parameters) + wanted_args.pop("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.can_read_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/tests/cogs/react_for_role/test_api.py b/tests/cogs/react_for_role/test_api.py index 5ee5ebf3..36dee8d9 100644 --- a/tests/cogs/react_for_role/test_api.py +++ b/tests/cogs/react_for_role/test_api.py @@ -1,12 +1,11 @@ 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 +from koala.cogs.react_for_role.api import RfrEndpoint, MESSAGE, REQUIRED_ROLES # Libs import discord @@ -16,20 +15,220 @@ @pytest.fixture -def api_client(bot: discord.ext.commands.Bot, aiohttp_client, loop ): +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_create_rfr(api_client): - resp = await api_client.post('/create', 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" -}) + +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 "rfr_id" in resp_json.keys() + 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 From 7c73d932a3fa73ad86fc632cd63d86d6de1edb50 Mon Sep 17 00:00:00 2001 From: JayDwee Date: Tue, 28 Mar 2023 21:49:15 +0100 Subject: [PATCH 23/25] temp: change dpytest requirement --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 20f4bb62..e607e9ed 100644 --- a/requirements.txt +++ b/requirements.txt @@ -6,7 +6,7 @@ certifi==2022.12.7 chardet==5.1.0 colorama==0.4.6 discord.py==2.2.2 -dpytest==0.6.3 +git+https://github.com/jaydwee/dpytest.git@master emoji==1.7.0 idna==3.4 iniconfig==2.0.0 From 247b93476df286574ad77b37cbe23b844c9f532f Mon Sep 17 00:00:00 2001 From: Jack Draper Date: Tue, 4 Apr 2023 20:19:53 +0100 Subject: [PATCH 24/25] chore: update dpytest version --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index e607e9ed..b7dcf800 100644 --- a/requirements.txt +++ b/requirements.txt @@ -6,7 +6,7 @@ certifi==2022.12.7 chardet==5.1.0 colorama==0.4.6 discord.py==2.2.2 -git+https://github.com/jaydwee/dpytest.git@master +dpytest==0.6.4 emoji==1.7.0 idna==3.4 iniconfig==2.0.0 From 74dd9ee45316138934fa23b6896d00f78004ab59 Mon Sep 17 00:00:00 2001 From: Jack Draper Date: Tue, 4 Apr 2023 20:35:13 +0100 Subject: [PATCH 25/25] test: fix failing tests --- tests/cogs/base/test_api.py | 14 ++++---------- 1 file changed, 4 insertions(+), 10 deletions(-) diff --git a/tests/cogs/base/test_api.py b/tests/cogs/base/test_api.py index 2eb6ef7b..c0a67842 100644 --- a/tests/cogs/base/test_api.py +++ b/tests/cogs/base/test_api.py @@ -232,12 +232,9 @@ async def test_post_load_cog_bad_req(api_client): 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' - }) + resp = await api_client.post('/load-cog', json={}) assert resp.status == BAD_REQUEST - assert await resp.text() == "400: Unsatisfied Arguments: {'package'}" + 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= @@ -286,12 +283,9 @@ async def test_post_unload_cog_not_loaded(api_client): 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' - }) + resp = await api_client.post('/unload-cog', json={}) assert resp.status == BAD_REQUEST - assert await resp.text() == "400: Unsatisfied Arguments: {'package'}" + 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=