diff --git a/ctfeed.py b/ctfeed.py index 330409e..56203b0 100644 --- a/ctfeed.py +++ b/ctfeed.py @@ -91,6 +91,7 @@ async def lifespan(app:FastAPI): app.include_router(router.user_router) app.include_router(router.ctf_router) app.include_router(router.config_router) +app.include_router(router.guild_router) # index @app.get("/", tags=["Shirakami Fubuki"]) diff --git a/notes/event.md b/notes/event.md index 76c5ccb..ef17dba 100644 --- a/notes/event.md +++ b/notes/event.md @@ -1,7 +1,98 @@ # About Event +## Rules - **時區都是 utc+0** (要記得轉換) - ``datetime.now(timezone.utc)`` - ``datetime_obj_with_timezone.astimezone(timezone.utc)`` -- 在讀取(即使用``read_event()``)的時候,如果不是確定「只會返回最多一個結果」的情境 - - finish_after=``int((datetime.now(timezone.utc) + timedelta(days=settings.DATABASE_SEARCH_DAYS)).timestamp())`` - - 我們需要限制讀出的數量,避免 DoS \ No newline at end of file +- 批量讀取 Event + - 使用 ``read_event_many()`` + - ``finish_after`` mode: ``finish_after=int((datetime.now(timezone.utc) + timedelta(days=settings.DATABASE_SEARCH_DAYS)).timestamp())`` + - 我們需要限制讀出的數量,避免 DoS + - 為避免日後讀取程式碼困難,使用不同的 mode 需要明確傳參(如 finish_before=None,就算是 None 也要傳) +- 請確保在操作 Database 中的 events table 時遵循以下流程,並確保整個流程被包覆在``try...except...finally...``中: + 1. 使用 ``src.crud.read_event(..., lock=True, duration=120) 對單個 event 加鎖,並獲取物件 + 2. (如果有需要,如創建頻道後將 ID 更新到資料庫)操作 Discord Bot + 3. 更新資料庫中的資料 + 4. 在``finally...``區塊中解鎖 + 5. 如有發生錯誤,在``except...``區塊中 rollback(例如:刪除創建出來的 Discord channel) + - 純讀取不受此限制 + +## Docs + +### ``read_event_one`` + +#### 用途 +- 用於讀取單一 Event(以 ``id`` 為主鍵) +- 可選擇是否同時嘗試加鎖(給後續更新流程使用) + +#### 使用方法 +- 只讀取(不加鎖) + - ``lock=False`` + - 回傳:``(event_db, None)`` +- 讀取並嘗試加鎖 + - ``lock=True`` 且 ``duration`` 必填 + - 成功:回傳 ``(event_db, lock_owner_token)`` + - Event 不存在:``NotFoundError`` + - Event 已被鎖住:``LockedError`` +- ``type`` 可為 ``ctftime`` / ``custom`` / ``None``,用來限制 Event 類型 +- ``archived`` 可為 ``True`` / ``False`` / ``None``,用來限制封存狀態 + +#### 設計說明 +- 單筆查詢一律以 ``Event.id`` 為核心條件 +- 加鎖模式採用原子條件更新(``locked_until`` + ``locked_by``)避免競態 +- 回傳的 ``lock_owner_token`` 需在後續 ``update_event`` / ``unlock_event`` 使用 + + +### ``read_event_many`` + +#### 用途 +- 用於讀取多筆 Event,給 API 列表查詢與背景工作使用 +- 支援 ``ctftime`` 與 ``custom`` 兩種查詢模式 + +#### 使用方法 +- ``type=ctftime`` + - ``finish_after`` mode + - 只傳 ``finish_after`` + - 不能同時傳 ``finish_before``、``limit``、``before_id`` + - ``finish_before`` mode + - ``finish_after`` 必須是 ``None`` + - ``limit`` 必填且需大於 0 + - 第一頁:``finish_before=None`` 且 ``before_id=None`` + - 下一頁:帶上前一頁最後一筆的 ``finish`` 和 ``id``(即 ``finish_before``、``before_id``) +- ``type=custom`` + - ``limit`` 必填且需大於 0 + - 第一頁:``before_id=None`` + - 下一頁:帶上前一頁最後一筆 ``id`` 到 ``before_id`` +- 不能傳 ``finish_after`` / ``finish_before`` +- ``archived`` 可選,用於限制封存狀態 + +#### 設計說明 +- ``ctftime`` 分頁採用複合游標條件: + - 排序:``ORDER BY finish DESC, id DESC`` + - 下一頁條件:``finish < finish_before`` 或 ``(finish == finish_before and id < before_id)`` +- ``custom`` 分頁採用 ``id`` 游標: + - 排序:``ORDER BY id DESC`` + - 下一頁條件:``id < before_id`` +- 參數組合會在函式內做嚴格檢查,不合法時拋 ``ValueError`` + + +### ``ctfmenu``(EventMenu / EventDetailMenu) + +#### 用途 +- Discord 互動式 Event 清單檢視(ctftime / custom) +- 提供分頁、切換 type、查看單筆 Event 詳細資訊 + +#### 設計說明 +- View timeout 設為 ``60s``,避免互動元件長時間掛著 +- ``ctftime`` 清單採「首次查詢後快取」: + - 第一次 ``build_embed_and_view()`` 查一次資料庫 + - 後續翻頁只吃記憶體快取,不重查 DB +- ``custom`` 清單採 cursor 分頁(``before_id`` + ``limit``): + - 每頁查 ``per_page + 1`` 判斷 ``has_next`` + - 使用 page cache(以頁碼快取已讀頁資料)避免回上一頁時被新資料擠動 +- ``EventDetailMenu`` 使用 ``read_event_one(lock=False, ...)`` 讀單筆,找不到時回傳 ``Event not found`` + +#### 常見坑 +- ``custom`` 分頁如果每次都重查 DB(不做頁面快取),新資料插入後會造成頁面漂移或看起來「有些 event 被擠掉」 +- ``ctftime`` 模式若每次翻頁都重查 DB,會有不必要的負擔;此場景改用快取較穩定 +- ``read_event_many`` 需明確傳 mode 參數(即使是 ``None`` 也傳)以提升可讀性並符合本專案規範 +- ``read_event_one`` / lock 流程內部使用 ``session.begin()``,caller 不要在外層再包 transaction 以避免巢狀交易風險 diff --git a/src/backend/channel_op.py b/src/backend/channel_op.py index 7611982..eb7ac03 100644 --- a/src/backend/channel_op.py +++ b/src/backend/channel_op.py @@ -1,4 +1,4 @@ -from typing import Optional, Dict, Any +from typing import Optional, Dict, Any, Tuple import logging from sqlalchemy.ext.asyncio import AsyncSession @@ -13,7 +13,7 @@ from src.utils import get_category from src.utils import ctf_api from src.utils import embed_creator -from src.bot import get_bot +from src.bot import get_guild from src import crud # channel_op = "event_op" @@ -21,8 +21,28 @@ # logging logger = logging.getLogger("uvicorn") +# utils +async def read_event_one_wrapper(session:AsyncSession, event_db_id:int) -> Tuple[model.Event, str]: + try: + event_db, lock_owner_token = await crud.read_event_one( + session=session, + lock=True, duration=120, + archived=False, # ensoure the Event isn't archived + id=event_db_id + ) + except crud.NotFoundError: + raise HTTPException(404, f"Event (id={event_db_id}) not found (archived, or invalid id)") + except crud.LockedError: + raise HTTPException(423, F"Event (id={event_db_id}) was locked. Try again later.") + except Exception as e: + logger.error(f"Can't get and lock Event (id={event_db_id}): {str(e)}") + raise HTTPException(500, f"Can't get and lock Event (id={event_db_id})") + + return event_db, lock_owner_token + + # functions -async def _create_channel(session:AsyncSession, member:discord.Member, event_db_id:int, lock_owner_token:str): +async def _create_channel(session:AsyncSession, member:discord.Member, event_db:model.Event, lock_owner_token:str) -> model.Event: # 在這個 function 有 exception 就直接 raise 出來 channel:Optional[discord.TextChannel] = None event_api:Optional[Dict[str, Any]] = None @@ -31,10 +51,7 @@ async def _create_channel(session:AsyncSession, member:discord.Member, event_db_ log_msg:str = "" # get guild - bot = await get_bot() - if (guild := bot.get_guild(settings.GUILD_ID)) is None: - logger.critical(f"Guild (id={settings.GUILD_ID}) not found") - raise HTTPException(500, f"Guild (id={settings.GUILD_ID}) not found") + guild = get_guild() # get category if (ctf_channel_category := get_category.get_category(guild, settings.CTF_CHANNEL_CATEGORY_ID)) is None: @@ -43,23 +60,13 @@ async def _create_channel(session:AsyncSession, member:discord.Member, event_db_ try: async with session.begin(): - # get a new event_db - events_db = await crud.read_event( - session, - id=event_db_id, - archived=False, # ensure the event isn't archived - lock_owner_token=lock_owner_token - ) - if len(events_db) != 1: - raise RuntimeError(f"Event (id={event_db_id}) not found") - event_db = events_db[0] ctftime_event = True if event_db.event_id is not None else False # check channel if (channel_id := event_db.channel_id) is not None and \ guild.get_channel(channel_id) is not None: # exists -> no need to create - return + return event_db if ctftime_event: events_api = await ctf_api.fetch_ctf_events(event_db.event_id) @@ -107,38 +114,24 @@ async def _create_channel(session:AsyncSession, member:discord.Member, event_db_ except Exception as e: logger.error(f"fail to send notification to channel (id={channel.id}): {str(e)}") # ignore exception - - return + return event_db -async def _join_channel(session:AsyncSession, member:discord.Member, event_db_id:int, lock_owner_token:str): + +async def _join_channel(session:AsyncSession, member:discord.Member, event_db:model.Event, lock_owner_token:str): # 在這個 function 有 exception 就直接 raise 出來 # get guild - bot = await get_bot() - if (guild := bot.get_guild(settings.GUILD_ID)) is None: - logger.critical(f"Guild (id={settings.GUILD_ID}) not found") - raise HTTPException(500, f"Guild (id={settings.GUILD_ID}) not found") + guild = get_guild() joined_channel = False # joined channel in Discord, but not in database joined = False # joined channel in Discord and database log_msg:str = "" try: async with session.begin(): - # get a new event_db - events_db = await crud.read_event( - session, - id=event_db_id, - archived=False, # ensure the Event isn't archived - lock_owner_token=lock_owner_token - ) - if len(events_db) != 1: - raise RuntimeError(f"Event (id={event_db_id}) not found") - event_db = events_db[0] - # check channel if (channel_id := event_db.channel_id) is None or \ (channel := guild.get_channel(channel_id)) is None: - raise RuntimeError(f"TextChannel for Event (id={event_db_id}) not found") + raise RuntimeError(f"TextChannel for Event (id={event_db.id}) not found") # join channel await channel.set_permissions(member, view_channel=True) @@ -146,10 +139,10 @@ async def _join_channel(session:AsyncSession, member:discord.Member, event_db_id # update database try: - await crud.join_event(session, event_db_id, member.id, lock_owner_token) + await crud.join_event(session, event_db.id, member.id, lock_owner_token) except IntegrityError: # ignore - raise HTTPException(409, f"The user (discord_id={member.id}) has joined the Event (id={event_db_id})") + raise HTTPException(409, f"The user (discord_id={member.id}) has joined the Event (id={event_db.id})") except Exception: raise joined = True @@ -196,22 +189,14 @@ async def create_and_join_channel(member:discord.Member, event_db_id:int): lock_owner_token:Optional[str] = None async with database.with_get_db() as session: # try to lock event - try: - lock_owner_token = await crud.try_lock_event(session, event_db_id, 120) - except crud.LockedError: - raise HTTPException(423, f"Event (id={event_db_id}) was locked. Try again later.") - except crud.NotFoundError: - raise HTTPException(404, f"Event (id={event_db_id}) not found") - except Exception as e: - logger.error(f"Can't lock Event (id={event_db_id}): {str(e)}") - raise HTTPException(f"Can't lock Event (id={event_db_id}): {str(e)}") + event_db, lock_owner_token = await read_event_one_wrapper(session, event_db_id) try: # try to create channel - await _create_channel(session, member, event_db_id, lock_owner_token) + event_db = await _create_channel(session, member, event_db, lock_owner_token) # join channel - await _join_channel(session, member, event_db_id, lock_owner_token) + await _join_channel(session, member, event_db, lock_owner_token) except Exception as e: if isinstance(e, HTTPException): raise @@ -239,10 +224,7 @@ async def archive_event(event_db_id:int, reason:str): event_db_returning = {} # get guild - bot = await get_bot() - if (guild := bot.get_guild(settings.GUILD_ID)) is None: - logger.critical(f"Guild (id={settings.GUILD_ID}) not found") - raise HTTPException(500, f"Guild (id={settings.GUILD_ID}) not found") + guild = get_guild() # get archive category if (archive_category := get_category.get_category(guild, settings.ARCHIVE_CATEGORY_ID)) is None: @@ -250,30 +232,9 @@ async def archive_event(event_db_id:int, reason:str): raise HTTPException(500, f"Archive Category (id={settings.ARCHIVE_CATEGORY_ID}) not found") async with database.with_get_db() as session: - # try to lock the Event - try: - lock_owner_token = await crud.try_lock_event(session, event_db_id, 120) - except crud.NotFoundError: - raise HTTPException(404, f"Event (id={event_db_id}) not found") - except crud.LockedError: - raise HTTPException(423, f"Event (id={event_db_id}) was locked. Try again later.") - except Exception as e: - logger.error(f"Can't lock Event (id={event_db_id}): {str(e)}") - raise HTTPException(500, f"Can't lock Event (id={event_db_id})") - + event_db, lock_owner_token = await read_event_one_wrapper(session, event_db_id) try: async with session.begin(): - # get a new event_db - events_db = await crud.read_event( - session=session, - archived=False, # ensure the Event isn't archived - id=event_db_id, - lock_owner_token=lock_owner_token, - ) - if len(events_db) != 1: - raise RuntimeError(f"Event (id={event_db_id}) not found") - event_db = events_db[0] - # update database event_db:model.Event = await crud.update_event( session=session, @@ -360,10 +321,7 @@ async def link_event_to_channel(event_db_id:int, channel_id:int): lock_owner_token = None # get guild - bot = await get_bot() - if (guild := bot.get_guild(settings.GUILD_ID)) is None: - logger.critical(f"Guild (id={settings.GUILD_ID}) not found") - raise HTTPException(500, f"Guild (id={settings.GUILD_ID}) not found") + guild = get_guild() # get channel if (channel := guild.get_channel(channel_id)) is None or \ @@ -371,30 +329,9 @@ async def link_event_to_channel(event_db_id:int, channel_id:int): raise HTTPException(400, f"Channel (id={channel_id}) not found") async with database.with_get_db() as session: - # try to lock the Event - try: - lock_owner_token = await crud.try_lock_event(session, event_db_id, 120) - except crud.NotFoundError: - raise HTTPException(404, f"Event (id={event_db_id}) not found") - except crud.LockedError: - raise HTTPException(423, f"Event (id={event_db_id}) was locked. Try again later.") - except Exception as e: - logger.error(f"Can't lock Event (id={event_db_id}): {str(e)}") - raise HTTPException(500, f"Can't lock Event (id={event_db_id})") - + event_db, lock_owner_token = await read_event_one_wrapper(session, event_db_id) try: async with session.begin(): - # get a new event_db - events_db = await crud.read_event( - session=session, - archived=False, # ensure the Event isn't archived - id=event_db_id, - lock_owner_token=lock_owner_token - ) - if len(events_db) != 1: - raise RuntimeError(f"Event (id={event_db_id}) not found") - event_db = events_db[0] - # update database event_db:model.Event = await crud.update_event( session=session, @@ -428,10 +365,7 @@ async def create_custom_event(title:str): :raise HTTPException: """ # get guild - bot = await get_bot() - if (guild := bot.get_guild(settings.GUILD_ID)) is None: - logger.critical(f"Guild (id={settings.GUILD_ID}) not found") - raise HTTPException(500, f"Guild (id={settings.GUILD_ID}) not found") + guild = get_guild() # create the custom event in database event_db_id = None diff --git a/src/backend/config.py b/src/backend/config.py index 1a4b94e..4201a74 100644 --- a/src/backend/config.py +++ b/src/backend/config.py @@ -7,6 +7,7 @@ from src import schema from src import crud +from src.bot import get_guild from src.database import model from src.database import database from src.config import settings, settings_lock @@ -49,7 +50,7 @@ async def check_config_valid_obj(guild:discord.Guild, key:str, value:Any) -> Tup return msg, _ -async def read_config(bot:commands.Bot, key:Optional[str]=None) -> schema.ConfigResponse: +async def read_config(key:Optional[str]=None) -> schema.ConfigResponse: """ Read config. @@ -61,10 +62,7 @@ async def read_config(bot:commands.Bot, key:Optional[str]=None) -> schema.Config :raise HTTPException: """ # get guild - guild = bot.get_guild(settings.GUILD_ID) - if guild is None: - logger.critical(f"Guild (id={settings.GUILD_ID}) not found") - raise HTTPException(500, f"Guild (id={settings.GUILD_ID}) not found") + guild = get_guild() # get config cache_config = {} @@ -123,7 +121,7 @@ async def update_config_cache(config:model.Config): logger.critical(f"fail to update cache of config (key={_k}) (maybe src.database.model.Config, src.database.model.config_info and src.config are out of sync): {str(e)}") -async def update_config(bot:commands.Bot, kv:Optional[Tuple]): +async def update_config(kv:Optional[Tuple]): """ Update Config in database and cache. @@ -133,10 +131,7 @@ async def update_config(bot:commands.Bot, kv:Optional[Tuple]): :raise HTTPException: """ # get guild - guild = bot.get_guild(settings.GUILD_ID) - if guild is None: - logger.critical(f"Guild (id={settings.GUILD_ID}) not found") - raise HTTPException(500, f"Guild (id={settings.GUILD_ID}) not found") + guild = get_guild() # check arguments arg = {} diff --git a/src/backend/event.py b/src/backend/event.py index 4a37ac1..0bf45ce 100644 --- a/src/backend/event.py +++ b/src/backend/event.py @@ -1,76 +1,26 @@ -from datetime import datetime, timezone, timedelta +from datetime import datetime, timezone from typing import Optional, List, Literal import logging -from sqlalchemy.ext.asyncio import AsyncSession -from fastapi import HTTPException import discord -from src.schema import Event, UserSimple, DiscordUser, DiscordChannel -from src.bot import get_bot -from src.config import settings -from src import crud +from src.database import model +from src.backend import security +from src import schema # logging logger = logging.getLogger("uvicorn") # functions -async def get_event( - session:AsyncSession, - # one - id:Optional[int]=None, - # many - type:Optional[Literal["ctftime", "custom"]]=None, - archived:Optional[bool]=None -) -> List[Event]: +async def format_event(guild:discord.Guild, events_db:List[model.Event]) -> List[schema.Event]: """ - Get Events. + :param guild: + :param events_db: - :param session: - :param id: - :param type: - :param archived: - - :return List[Event]: - - :raise HTTPException: + :return List[schema.Event]: A list of formatted Events. """ - # get guild - bot = await get_bot() - if (guild := bot.get_guild(settings.GUILD_ID)) is None: - logger.critical(f"Guild (id={settings.GUILD_ID}) not found") - raise HTTPException(500, f"Guild (id={settings.GUILD_ID}) not found") - - # build arguments - finish_after = int((datetime.now(timezone.utc) + timedelta(days=settings.DATABASE_SEARCH_DAYS)).timestamp()) - - # get events from database - try: - async with session.begin(): - if id is not None: - # one - events_db = await crud.read_event(session=session, id=id) - if len(events_db) != 1: - raise HTTPException(404, f"Event (id={id}) not found") - else: - # many - if type not in ["ctftime", "custom"]: - raise HTTPException(400, "type should be ctftime or custom when id is None") - - events_db = await crud.read_event( - session=session, - type=type, - archived=archived, - finish_after=(finish_after if type == "ctftime" else None) - ) - except Exception as e: - if isinstance(e, HTTPException): - raise - logger.error(f"fail to read Events from database: {str(e)}") - raise HTTPException(500, "fail to read Events from database") - - result:List[Event] = [] + result:List[schema.Event] = [] for event in events_db: # event attributes event_type:Literal["ctftime", "custom"] = "ctftime" if event.event_id is not None else "custom" @@ -84,37 +34,41 @@ async def get_event( now_running = False # channel - channel:Optional[DiscordChannel] = None + channel:Optional[schema.DiscordTextChannel] = None if event.channel_id is not None: discord_channel = guild.get_channel(event.channel_id) if isinstance(discord_channel, discord.TextChannel): - channel = DiscordChannel( + channel = schema.DiscordTextChannel( id=discord_channel.id, jump_url=discord_channel.jump_url, name=discord_channel.name ) # users - users:List[UserSimple] = [] + users:List[schema.UserSimple] = [] for db_user in event.users: - discord_user:Optional[DiscordUser] = None + discord_user:Optional[schema.DiscordUser] = None + user_role:List[schema.UserRole] = [] member = guild.get_member(db_user.discord_id) if member is not None: - discord_user = DiscordUser( + discord_user = schema.DiscordUser( display_name=member.display_name, id=member.id, name=member.name ) + + user_role = await security.get_role(member) - users.append(UserSimple( + users.append(schema.UserSimple( discord_id=db_user.discord_id, + user_role=user_role, status=db_user.status, skills=db_user.skills, rhythm_games=db_user.rhythm_games, discord=discord_user )) - result.append(Event( + result.append(schema.Event( id=event.id, archived=event.archived, event_id=event.event_id, diff --git a/src/backend/security.py b/src/backend/security.py index 0fdfd10..cf0f147 100644 --- a/src/backend/security.py +++ b/src/backend/security.py @@ -1,4 +1,4 @@ -from typing import Optional +from typing import Optional, List import logging from sqlalchemy.exc import IntegrityError @@ -8,12 +8,32 @@ from src.config import settings, settings_lock from src.database.database import with_get_db -from src.bot import get_bot +from src.bot import get_guild +from src import schema from src import crud # logging logger = logging.getLogger("uvicorn") +# utils +async def get_role(member:discord.Member) -> List[schema.UserRole]: + roles = [] + if member.guild_permissions.administrator == True: + roles.append(schema.UserRole.administrator) + + async with settings_lock: + pm_role_id = settings.PM_ROLE_ID + member_role_id = settings.MEMBER_ROLE_ID + + if member.get_role(pm_role_id): + roles.append(schema.UserRole.pm) + + if member.get_role(member_role_id): + roles.append(schema.UserRole.member) + + return roles + + # functions async def check_administrator(discord_id:int) -> Optional[discord.Member]: """ @@ -24,12 +44,10 @@ async def check_administrator(discord_id:int) -> Optional[discord.Member]: :param discord_id: :return Optional[discord.Member]: + + :raise HTTPException: """ - bot = await get_bot() - guild = bot.get_guild(settings.GUILD_ID) - if guild is None: - logger.critical(f"Guild (id={settings.GUILD_ID}) not found") - return None + guild = get_guild() # check whether the user is in the guild member = guild.get_member(discord_id) @@ -56,13 +74,8 @@ async def check_user(discord_id:int, force_pm:bool) -> discord.Member: :raise HTTPException: """ - bot:commands.Bot = await get_bot() - # get guild - guild = bot.get_guild(settings.GUILD_ID) - if guild is None: - logger.critical(f"Guild (id={settings.GUILD_ID}) not found") - raise HTTPException(500, f"Guild (id={settings.GUILD_ID}) not found") + guild = get_guild() # get member member = guild.get_member(discord_id) @@ -132,7 +145,13 @@ async def check_user_and_auto_register( # for discord async def discord_check_administrator(interaction:discord.Interaction) -> bool: - if (await check_administrator(interaction.user.id)) is None: + try: + member = await check_administrator(interaction.user.id) + except HTTPException: + await interaction.response.send_message(f"Guild (id={settings.GUILD_ID}) not found", ephemeral=True) + return False + + if member is None: await interaction.response.send_message("Forbidden", ephemeral=True) return False return True diff --git a/src/backend/user.py b/src/backend/user.py index db72675..7787bad 100644 --- a/src/backend/user.py +++ b/src/backend/user.py @@ -7,8 +7,8 @@ from src.database import model from src.backend import security -from src.schema import User, DiscordUser, EventSimple -from src.bot import get_bot +from src.schema import UserRole, User, DiscordUser, EventSimple +from src.bot import get_guild from src.config import settings from src import crud @@ -28,10 +28,7 @@ async def get_user(session:AsyncSession, discord_id:Optional[int]=None) -> List[ :raise HTTPException: """ # get guild - bot = await get_bot() - if (guild := bot.get_guild(settings.GUILD_ID)) is None: - logger.critical(f"Guild (id={settings.GUILD_ID}) not found") - raise HTTPException(500, f"Guild (id={settings.GUILD_ID}) not found") + guild = get_guild() # get user from database try: @@ -79,15 +76,19 @@ async def get_user(session:AsyncSession, discord_id:Optional[int]=None) -> List[ # discord user discord_user:Optional[DiscordUser] = None + user_role:List[UserRole] = [] if member is not None: discord_user = DiscordUser( display_name=member.display_name, id=member.id, name=member.name ) + + user_role = await security.get_role(member) result.append(User( discord_id=db_user.discord_id, + user_role=user_role, status=db_user.status, skills=db_user.skills, rhythm_games=db_user.rhythm_games, diff --git a/src/bgtask/detect_event_update_and_remove.py b/src/bgtask/detect_event_update_and_remove.py index 9cb0aeb..47a7afc 100644 --- a/src/bgtask/detect_event_update_and_remove.py +++ b/src/bgtask/detect_event_update_and_remove.py @@ -7,7 +7,7 @@ from src.utils import embed_creator from src.utils import notification from src.config import settings -from src.bot import get_bot +from src.bot import get_guild from src.backend import channel_op from src import crud @@ -27,41 +27,34 @@ async def check_and_update_event(event_db_id:int, event_api:Dict[str, Any]): event_db_returning = {"updated": False} # get guild - bot = await get_bot() - guild = bot.get_guild(settings.GUILD_ID) - if guild is None: - logger.critical(f"Guild (id={settings.GUILD_ID}) not found") + try: + guild = get_guild() + except Exception: return # update database async with database.with_get_db() as session: - # try to lock the Event + # get a new event_db try: - lock_owner_token = await crud.try_lock_event(session, event_db_id, 120) + event_db, lock_owner_token = await crud.read_event_one( + session=session, + lock=True, duration=120, + type="ctftime", + archived=False, # ensoure the event isn't archived + id=event_db_id + ) except crud.NotFoundError: - logger.error(f"Event (id={event_db_id}) not found.") + logger.warning(f"Event (id={event_db_id}) not found.") return except crud.LockedError: - logger.info(f"Event (id={event_db_id}) was locked. Skipped...") + logger.warning(f"Event (id={event_db_id}) was locked. Skipped...") return except Exception as e: - logger.error(f"Can't lock Event (id={event_db_id}): {str(e)}") + logger.error(f"Can't get and lock Event (id={event_db_id}): {str(e)}") return - + try: async with session.begin(): - # get a new event_db - events_db = await crud.read_event( - session=session, - type="ctftime", - archived=False, # ensure the event isn't archived - id=event_db_id, - lock_owner_token=lock_owner_token - ) - if len(events_db) != 1: - raise RuntimeError(f"Event (id={event_db_id}) not found") - event_db = events_db[0] - # check if event_db.title != ntitle or \ event_db.start != int(nstart.timestamp()) or \ @@ -149,11 +142,14 @@ async def _detect_event_update_and_remove(): # get all non-archived CTFTime events from database try: async with database.with_get_db() as session: - events_db = await crud.read_event( + events_db = await crud.read_event_many( session=session, type="ctftime", archived=False, - finish_after=int((datetime.now(timezone.utc) + timedelta(days=settings.DATABASE_SEARCH_DAYS)).timestamp()) + limit=None, + finish_after=int((datetime.now(timezone.utc) + timedelta(days=settings.DATABASE_SEARCH_DAYS)).timestamp()), + finish_before=None, + before_id=None ) events_db_returning = [ diff --git a/src/bgtask/detect_events_new.py b/src/bgtask/detect_events_new.py index c8c7446..100eac9 100644 --- a/src/bgtask/detect_events_new.py +++ b/src/bgtask/detect_events_new.py @@ -30,11 +30,14 @@ async def _detect_events_new(): # archived=None - get both archived and non-archived events to avoid missing any events try: async with database.with_get_db() as session: - events_db = await crud.read_event( + events_db = await crud.read_event_many( session, type="ctftime", archived=None, - finish_after=int((datetime.now(timezone.utc)+timedelta(days=settings.DATABASE_SEARCH_DAYS)).timestamp()) + limit=None, + finish_after=int((datetime.now(timezone.utc)+timedelta(days=settings.DATABASE_SEARCH_DAYS)).timestamp()), + finish_before=None, + before_id=None ) events_db_event_id = [event.event_id for event in events_db] except Exception as e: @@ -85,4 +88,4 @@ async def _detect_events_new(): logger.error(f"fail to send notification to announcement channel: {str(e)}") # ignore exception - return \ No newline at end of file + return diff --git a/src/bgtask/recover_scheduled_events.py b/src/bgtask/recover_scheduled_events.py index 7840735..2b6694c 100644 --- a/src/bgtask/recover_scheduled_events.py +++ b/src/bgtask/recover_scheduled_events.py @@ -5,7 +5,7 @@ import discord from src.database import database -from src.bot import get_bot +from src.bot import get_guild from src.config import settings from src import crud @@ -15,42 +15,36 @@ # functions async def do_recover(event_db_id:int): # get guild - bot = await get_bot() - if (guild := bot.get_guild(settings.GUILD_ID)) is None: - logger.critical(f"Guild (id={settings.GUILD_ID}) not found") + try: + guild = get_guild() + except Exception: return lock_owner_token:Optional[str] = None sc:Optional[discord.ScheduledEvent] = None need_create = False async with database.with_get_db() as session: - # try to lock the Event + # get a new event_db try: - lock_owner_token = await crud.try_lock_event(session, event_db_id, 120) + event_db, lock_owner_token = await crud.read_event_one( + session=session, + lock=True, duration=120, + type="ctftime", + archived=False, # ensure the event isn't archived + id=event_db_id, + ) except crud.NotFoundError: - logger.error(f"Event (id={event_db_id}) not found.") + logger.warning(f"Event (id={event_db_id}) not found.") return except crud.LockedError: - logger.info(f"Event (id={event_db_id}) was locked. Skipped...") + logger.warning(f"Event (id={event_db_id}) was locked. Skipped...") return except Exception as e: - logger.error(f"Can't lock Event (id={event_db_id}): {str(e)}") + logger.error(f"Can't get and lock Event (id={event_db_id}): {str(e)}") return try: async with session.begin(): - # get a new event_db - events_db = await crud.read_event( - session=session, - type="ctftime", - archived=False, # ensure the event isn't archived - id=event_db_id, - lock_owner_token=lock_owner_token - ) - if len(events_db) != 1: - raise RuntimeError(f"Event (id={event_db_id}) not found") - event_db = events_db[0] - # check if event_db.channel_id is None: # The event doesn't have a channel, so no need to create or update it's scheduled event @@ -133,11 +127,14 @@ async def _recover_scheduled_events(): """ async with database.with_get_db() as session: try: - events_db = await crud.read_event( + events_db = await crud.read_event_many( session, type="ctftime", archived=False, - finish_after=int((datetime.now(timezone.utc) + timedelta(days=settings.DATABASE_SEARCH_DAYS)).timestamp()) + limit=None, + finish_after=int((datetime.now(timezone.utc) + timedelta(days=settings.DATABASE_SEARCH_DAYS)).timestamp()), + finish_before=None, + before_id=None ) except Exception as e: logger.error(f"fail to get known CTF events from database: {str(e)}") @@ -147,4 +144,4 @@ async def _recover_scheduled_events(): for event_db_id in events_db_id: await do_recover(event_db_id) - return \ No newline at end of file + return diff --git a/src/bot.py b/src/bot.py index d666109..643c8ca 100644 --- a/src/bot.py +++ b/src/bot.py @@ -2,7 +2,9 @@ import asyncio import glob import pathlib +import traceback +from fastapi import HTTPException from discord.ext import commands import discord @@ -40,6 +42,12 @@ def load_cogs(): logger.critical(f"fail to load {extension_name}: {str(e)}") +# global error handler +@bot.event +async def on_error(event: str, *args, **kwargs): + logger.error(f"Unhandled exception in event ({event}): {''.join(traceback.format_exc())}") + + # startup and shutdown async def main(): load_cogs() @@ -61,5 +69,20 @@ async def stop_bot(): # get bot -async def get_bot() -> commands.Bot: +def get_bot() -> commands.Bot: return bot + + +# get guild +def get_guild() -> discord.Guild: + """ + Get the guild. + + :return discord.Guild: + + :raise HTTPException: + """ + if (guild := bot.get_guild(settings.GUILD_ID)) is None: + logger.critical(f"Guild (id={settings.GUILD_ID}) not found") + raise HTTPException(500, f"Guild (id={settings.GUILD_ID}) not found") + return guild \ No newline at end of file diff --git a/src/cog/config.py b/src/cog/config.py index 8878d8a..4b17362 100644 --- a/src/cog/config.py +++ b/src/cog/config.py @@ -24,7 +24,7 @@ async def build_embed_and_view(self) -> discord.Embed: # build embed try: - config_info = await config.read_config(self.bot, self.state if self.state != "MAIN" else None) + config_info = await config.read_config(self.state if self.state != "MAIN" else None) except Exception as e: return discord.Embed(title=f"Fail to read config", description=str(e), color=discord.Color.red()) @@ -81,7 +81,7 @@ async def _build_view(self): async def on_change_page(self, interaction:discord.Interaction): # check permission - if not(await security.discord_check_administrator(interaction)): + if not (await security.discord_check_administrator(interaction)): return # check argument @@ -125,7 +125,7 @@ async def on_edit(self, interaction:discord.Interaction): # update try: - await config.update_config(self.bot, (self.state, value)) + await config.update_config((self.state, value)) except Exception as e: await interaction.response.send_message(f"fail to update config (key={self.state}): {str(e)}", ephemeral=True) return diff --git a/src/cog/ctfmenu.py b/src/cog/ctfmenu.py index ada1eec..f3fa935 100644 --- a/src/cog/ctfmenu.py +++ b/src/cog/ctfmenu.py @@ -1,5 +1,5 @@ from datetime import datetime, timezone, timedelta -from typing import Literal, Optional, List +from typing import Literal, Optional, List, Dict import logging import math @@ -31,7 +31,7 @@ def _format_channel_info(guild: Optional[discord.Guild], channel_id: Optional[in # views class EventMenu(discord.ui.View): def __init__(self, bot: commands.Bot, owner_id: int, type: Literal["ctftime", "custom"]): - super().__init__(timeout=None) + super().__init__(timeout=60) self.bot = bot self.owner_id = owner_id self.type = type @@ -39,6 +39,14 @@ def __init__(self, bot: commands.Bot, owner_id: int, type: Literal["ctftime", "c self.per_page = 5 self.events: List[model.Event] = [] + self.ctftime_events_cache: List[model.Event] = [] + self.ctftime_cache_ready = False + + self.custom_before_id_history: List[Optional[int]] = [None] + self.custom_has_next = False + self.custom_page_cache: Dict[int, List[model.Event]] = {} + self.custom_page_has_next_cache: Dict[int, bool] = {} + async def _check_permission(self, interaction: discord.Interaction) -> Optional[discord.Member]: if (member := (await security.discord_check_user_and_auto_register(interaction, False))) is None: @@ -49,13 +57,20 @@ async def _check_permission(self, interaction: discord.Interaction) -> Optional[ return member - async def _refresh_view(self, total_pages: int): - self.prev_page.disabled = self.page <= 0 - self.next_page.disabled = self.page >= total_pages - 1 - - start = self.page * self.per_page - end = start + self.per_page - current = self.events[start:end] + async def _refresh_view(self, total_pages: Optional[int] = None): + if self.type == "ctftime": + self.prev_page.disabled = self.page <= 0 + if total_pages is None: + raise RuntimeError("total_pages should not be None for ctftime menu") + self.next_page.disabled = self.page >= total_pages - 1 + + start = self.page * self.per_page + end = start + self.per_page + current = self.events[start:end] + else: + self.prev_page.disabled = self.page <= 0 + self.next_page.disabled = self.custom_has_next is False + current = self.events self.select_event.disabled = len(current) == 0 if len(current) == 0: @@ -85,39 +100,68 @@ async def _refresh_view(self, total_pages: int): async def build_embed_and_view(self) -> discord.Embed: try: async with database.with_get_db() as session: - finish_after = None if self.type == "ctftime": - finish_after = int((datetime.now(timezone.utc) + timedelta(days=settings.DATABASE_SEARCH_DAYS)).timestamp()) - - self.events = await crud.read_event( - session=session, - type=self.type, - archived=False, - finish_after=finish_after - ) + if self.ctftime_cache_ready is False: + self.ctftime_events_cache = await crud.read_event_many( + session=session, + type="ctftime", + archived=False, + limit=None, + finish_after=int((datetime.now(timezone.utc) + timedelta(days=settings.DATABASE_SEARCH_DAYS)).timestamp()), + finish_before=None, + before_id=None, + ) + self.ctftime_cache_ready = True + + self.events = self.ctftime_events_cache + else: + if self.page < 0: + self.page = 0 + if self.page >= len(self.custom_before_id_history): + self.page = len(self.custom_before_id_history) - 1 + + if self.page in self.custom_page_cache: + self.events = self.custom_page_cache[self.page] + self.custom_has_next = self.custom_page_has_next_cache[self.page] + else: + before_id = self.custom_before_id_history[self.page] + events = await crud.read_event_many( + session=session, + type="custom", + archived=False, + limit=self.per_page + 1, + finish_after=None, + finish_before=None, + before_id=before_id, + ) + self.custom_has_next = len(events) > self.per_page + self.events = events[:self.per_page] + self.custom_page_cache[self.page] = self.events + self.custom_page_has_next_cache[self.page] = self.custom_has_next except Exception as e: logger.error(f"fail to read Events: {str(e)}") return discord.Embed(title="Fail to read events", color=discord.Color.red()) - if self.type == "custom": - self.events = sorted(self.events, key=lambda e: e.id) - - total_pages = max(1, math.ceil(len(self.events) / self.per_page)) - if self.page >= total_pages: - self.page = total_pages - 1 - if self.page < 0: - self.page = 0 - - start = self.page * self.per_page - end = start + self.per_page - current = self.events[start:end] + if self.type == "ctftime": + total_pages = max(1, math.ceil(len(self.events) / self.per_page)) + if self.page >= total_pages: + self.page = total_pages - 1 + if self.page < 0: + self.page = 0 + + start = self.page * self.per_page + end = start + self.per_page + current = self.events[start:end] + else: + current = self.events + display_start = self.page * self.per_page title = "CTFTime Events" if self.type == "ctftime" else "Custom Events" if len(current) == 0: description = "(No events)" else: lines = [] - for idx, e in enumerate(current, start=start + 1): + for idx, e in enumerate(current, start=display_start + 1): channel_created = "[⭐️ Channel created]" if e.channel_id is not None else "" if self.type == "ctftime": @@ -138,9 +182,13 @@ async def build_embed_and_view(self) -> discord.Embed: description = "\n".join(lines) embed = discord.Embed(title=title, description=description, color=discord.Color.green()) - embed.set_footer(text=f"Page {self.page + 1}/{total_pages} | Total {len(self.events)}") + if self.type == "ctftime": + embed.set_footer(text=f"Page {self.page + 1}/{total_pages} | Total {len(self.events)}") + await self._refresh_view(total_pages) + else: + embed.set_footer(text=f"Page {self.page + 1}") + await self._refresh_view() - await self._refresh_view(total_pages) return embed @@ -148,6 +196,10 @@ async def build_embed_and_view(self) -> discord.Embed: async def prev_page(self, button: discord.ui.Button, interaction: discord.Interaction): if await self._check_permission(interaction) is None: return + if self.page <= 0: + await interaction.response.edit_message(view=self) + return + self.page -= 1 embed = await self.build_embed_and_view() await interaction.response.edit_message(embed=embed, view=self) @@ -157,6 +209,14 @@ async def prev_page(self, button: discord.ui.Button, interaction: discord.Intera async def next_page(self, button: discord.ui.Button, interaction: discord.Interaction): if await self._check_permission(interaction) is None: return + if self.type == "custom": + if len(self.events) == 0 or self.custom_has_next is False: + await interaction.response.edit_message(view=self) + return + + if self.page + 1 >= len(self.custom_before_id_history): + self.custom_before_id_history.append(self.events[-1].id) + self.page += 1 embed = await self.build_embed_and_view() await interaction.response.edit_message(embed=embed, view=self) @@ -211,7 +271,7 @@ async def create_custom_event(self, button: discord.ui.Button, interaction: disc class EventDetailMenu(discord.ui.View): def __init__(self, bot: commands.Bot, owner_id: int, event_db_id: int, type: Literal["ctftime", "custom"]): - super().__init__(timeout=None) + super().__init__(timeout=60) self.bot = bot self.owner_id = owner_id self.event_db_id = event_db_id @@ -230,15 +290,16 @@ async def _check_permission(self, interaction: discord.Interaction, force_pm:boo async def _read_event(self) -> Optional[model.Event]: try: async with database.with_get_db() as session: - events = await crud.read_event( + event_db, _ = await crud.read_event_one( session=session, + lock=False, type=self.type, archived=False, id=self.event_db_id ) - if len(events) != 1: - return None - return events[0] + return event_db + except crud.NotFoundError: + return None except Exception as e: logger.error(f"fail to read Event (id={self.event_db_id}): {str(e)}") return None diff --git a/src/cog/user.py b/src/cog/user.py index 815e689..ba82570 100644 --- a/src/cog/user.py +++ b/src/cog/user.py @@ -33,17 +33,18 @@ async def build_embed_and_view(self) -> discord.Embed: # build embed color = discord.Color.green() if user_s.status == model.Status.online else discord.Color.red() + user_role_s = ", ".join([r.value for r in user_s.user_role]) if user_s.discord is not None: embed = discord.Embed( title=f"{user_s.discord.display_name}", - description=f"Name: {user_s.discord.name}\nID: {user_s.discord.id}", + description=f"Name: {user_s.discord.name}\nID: {user_s.discord.id}\nRoles: {user_role_s}", color=color ) else: embed = discord.Embed( title=f"(Invalid)", - description=f"Name: (Invalid)\nID: {user_s.discord_id}", + description=f"Name: (Invalid)\nID: {user_s.discord_id}\nRoles: {user_role_s}", color=color ) diff --git a/src/crud/__init__.py b/src/crud/__init__.py index 09409d2..8aff394 100644 --- a/src/crud/__init__.py +++ b/src/crud/__init__.py @@ -1,8 +1,8 @@ from .config import create_or_update_config, read_config from .user import create_user, read_user, update_user from .event import ( - try_lock_event, unlock_event, + unlock_event, NotFoundError, LockedError, join_event, delete_user_in_event, - create_event, read_event, read_ctfime_events_need_archive, update_event, + create_event, read_event_one, read_event_many, read_ctfime_events_need_archive, update_event, ) \ No newline at end of file diff --git a/src/crud/config.py b/src/crud/config.py index f7be554..086b222 100644 --- a/src/crud/config.py +++ b/src/crud/config.py @@ -59,6 +59,7 @@ async def create_or_update_config( try: result = (await session.execute(stmt)).scalar_one() await session.flush() + await session.refresh(result) return result except Exception: raise diff --git a/src/crud/event.py b/src/crud/event.py index 09434dc..1c72ff8 100644 --- a/src/crud/event.py +++ b/src/crud/event.py @@ -1,4 +1,4 @@ -from typing import List, Optional, Literal +from typing import List, Optional, Literal, Tuple, Union from datetime import datetime, timedelta, timezone import hashlib import os @@ -17,72 +17,149 @@ class NotFoundError(Exception): class LockedError(Exception): pass -async def try_lock_event(session:AsyncSession, id:int, duration:int) -> str: - """ - Try to lock an Event. - - :param session: - :param id: The Event which you want to lock. - :param duration: How long you want to lock the Event (in seconds). - - :return str: Lock owner token. - - :raise NotFoundError: Can't find the Event. - :raise LockedError: The Event was locked. - :raise (Exception from sqlalchemy): - """ - # time and lock_owner_token - time_now = datetime.now(timezone.utc) - locked_until = time_now + timedelta(seconds=duration) - lock_owner_token = hashlib.sha256(os.urandom(32)).hexdigest() - - # stmt - check_exists_cte = ( - sqlalchemy.select(Event.id) \ - .where(Event.id == id) - ).cte("check_exists_cte") - - try_lock_cte = ( - sqlalchemy.update(Event) \ - .where(Event.id == id) \ - .where(sqlalchemy.or_( - Event.locked_until == None, - Event.locked_until < int(time_now.timestamp()) # expired - )) \ - .values( - locked_until=int(locked_until.timestamp()), - locked_by=lock_owner_token - ) \ - .returning(Event) - ).cte("try_lock_cte") - - check_lock = sqlalchemy.exists( - sqlalchemy.select(try_lock_cte.c.id) \ - .where(try_lock_cte.c.id == id) \ - .where(try_lock_cte.c.locked_by == lock_owner_token) - ) - - stmt = sqlalchemy.select( - sqlalchemy.case( - (check_lock, "success"), - else_="locked" - ), - ) \ - .where(check_exists_cte.c.id == id) - - # execute - async with session.begin(): - status = (await session.execute(stmt)).one_or_none() - - if status is None: - raise NotFoundError - else: - status = status[0] - if status == "success": - return lock_owner_token - elif status == "locked": - raise LockedError - +#async def read_event( +# session:AsyncSession, +# # lock +# lock:bool, +# duration:Optional[int]=None, +# # filters +# type:Optional[Literal["ctftime", "custom"]]=None, +# archived:Optional[bool]=None, +# id:Optional[int]=None, +# channel_id:Optional[int]=None, +# # only for CTFTime Events +# event_id:Optional[int]=None, +# finish_after:Optional[int]=None, +#) -> Tuple[List[Event], Optional[str]]: +# """ +# Read Events and try to lock an Event. +# +# :param session: +# :param lock: Whether to lock the Event. +# :param duration: +# :param type: Search "ctftime", "custom" Events, or ``None`` to search both types of Events. +# :param archived: Search archived, non-archived Events, or ``None`` to search both types of Events. +# :param id: +# :param channel_id: Discord channel ID. +# :param event_id: (CTFTime Event only) CTFTime Event id. +# :param finish_after: (CTFTime Event only) search CTFTime Events which finish after ``finish_after``. +# +# :return List[Event]: A list of Events. +# :return str: Lock owner token. +# +# :raise NotFoundError: Can't find the Event (when id, channel_id or event_id is not None). +# :raise LockedError: The Event was locked. +# :raise ValueError: +# :raise (Exception from sqlalchemy): +# """ +# # functions +# def _build_filter(stmt): +# if type is not None: +# if type == "ctftime": +# if event_id is not None: +# stmt = stmt.where(Event.event_id == event_id) +# else: +# stmt = stmt.where(Event.event_id != None) +# +# if finish_after is not None: +# stmt = stmt.where(Event.finish >= finish_after) +# +# # sorting +# if isinstance(stmt, sqlalchemy.Select): +# stmt = stmt.order_by(sqlalchemy.asc(Event.finish)) +# elif type == "custom": +# stmt = stmt.where(Event.event_id == None) +# else: +# raise ValueError(f"type should be \"ctftime\", \"custom\" or None") +# +# if archived is not None: +# stmt = stmt.where(Event.archived == archived) +# +# if id is not None: +# stmt = stmt.where(Event.id == id) +# +# if channel_id is not None: +# stmt = stmt.where(Event.channel_id == channel_id) +# +# return stmt +# +# # check exists stmt +# check_exists = sqlalchemy.select(Event) \ +# .options(selectinload(Event.users)) +# check_exists:sqlalchemy.Select = _build_filter(check_exists) +# +# # no need to lock -> execute and return +# if lock == False: +# try: +# results = (await session.execute(check_exists)).scalars().all() +# except Exception: +# raise +# if len(results) == 0: +# raise NotFoundError +# return results, None +# +# # need to lock +# # only effective when id is not None +# if id is None: +# raise ValueError("id should not be None when lock is True") +# +# # argument check +# if duration is None: +# raise ValueError("duration should not be None when lock is True") +# +# # prepare arguments +# time_now = datetime.now(timezone.utc) +# locked_until = time_now + timedelta(seconds=duration) +# lock_owner_token = hashlib.sha256(os.urandom(32)).hexdigest() +# +# # stmt +# check_exists_cte = check_exists.cte("check_exists_cte") +# +# try_lock = sqlalchemy.update(Event) +# try_lock:sqlalchemy.Update = _build_filter(try_lock) +# try_lock_cte = ( +# try_lock +# .where(sqlalchemy.or_( +# Event.locked_until == None, +# Event.locked_until < int(time_now.timestamp()) +# )) +# .values( +# locked_until=int(locked_until.timestamp()), +# locked_by=lock_owner_token +# ) +# .returning(Event) +# ).cte("try_lock_cte") +# +# check_lock = sqlalchemy.exists( +# sqlalchemy.select(try_lock_cte.c.id) \ +# .where(try_lock_cte.c.id == id) \ +# .where(try_lock_cte.c.locked_by == lock_owner_token) +# ) +# +# stmt = sqlalchemy.select( +# sqlalchemy.case( +# (check_lock, "success"), +# else_="locked" +# ), +# Event +# ) \ +# .options(selectinload(Event.users)) \ +# .join(check_exists_cte, check_exists_cte.c.id == Event.id) \ +# .where(check_exists_cte.c.id == id) +# +# # execute +# async with session.begin(): +# results = (await session.execute(stmt)).all() +# if len(results) == 0: +# raise NotFoundError +# else: +# status = results[0][0] +# event_db = results[0][1] +# if status == "success": +# return [event_db], lock_owner_token +# elif status == "locked": +# raise LockedError +# async def unlock_event(session:AsyncSession, id:int, lock_owner_token:str) -> bool: """ @@ -153,7 +230,7 @@ async def join_event( # execute try: - result = (await session.execute(stmt)).one() + (await session.execute(stmt)).one() await session.flush() return except Exception: @@ -274,41 +351,174 @@ async def create_event( try: result = (await session.execute(stmt)).scalar_one() await session.flush() + await session.refresh(result) return result except Exception: raise # read -async def read_event( +async def read_event_one( session:AsyncSession, - # lock - lock_owner_token:Optional[str]=None, - # conditions + id:int, + lock:bool, + duration:Optional[int]=None, type:Optional[Literal["ctftime", "custom"]]=None, archived:Optional[bool]=None, - id:Optional[int]=None, - channel_id:Optional[int]=None, - # only for CTFTime Events - event_id:Optional[int]=None, +) -> Tuple[Event, Optional[str]]: + """ + Read one Event and try to lock an Event (if you want). + + Inside this function, it uses ``async with session.begin()``. + + :param session: + :param id: + :param lock: Whether to lock the Event. + :param duration: How long you want to lock the Event (in seconds). + :param type: Search ``ctftime``, ``custom`` Events, or ``None`` to search both types of Events. + :param archived: Search archived, non-archived Events, or ``None`` to search both types of Events. + + :return Event: + :return Optional[str]: Lock owner token. + + :raise NotFoundError: Can't find the Event. + :raise LockedError: The Event was locked. + :raise ValueError: + :raise RuntimeError: + :raise (Exception from sqlalchemy): + """ + # functions + def _build_filter(stmt:Union[sqlalchemy.Select, sqlalchemy.Update]) -> Union[sqlalchemy.Select, sqlalchemy.Update]: + stmt = stmt.where(Event.id == id) + + if type is not None: + if type == "ctftime": + stmt = stmt.where(Event.event_id != None) + elif type == "custom": + stmt = stmt.where(Event.event_id == None) + else: + raise ValueError(f"type should be \"ctftime\", \"custom\" or None") + + if archived is not None: + stmt = stmt.where(Event.archived == archived) + + return stmt + + # check exists stmt + check_exists:sqlalchemy.Select = _build_filter(sqlalchemy.select(Event)) + + # no need to lock -> execute and return + if lock == False: + check_exists = check_exists.options(selectinload(Event.users)) + async with session.begin(): + try: + event_db = (await session.execute(check_exists)).scalar_one_or_none() + except Exception: + raise + if event_db is None: + raise NotFoundError + return event_db, None + + # need to lock + # argument check + if duration is None: + raise ValueError("duration should not be None when lock is True") + + # prepare arguments + time_now = datetime.now(timezone.utc) + locked_until = time_now + timedelta(seconds=duration) + lock_owner_token = hashlib.sha256(os.urandom(32)).hexdigest() + + # stmt + check_exists_cte = check_exists.cte("check_exists_cte") + + try_lock:sqlalchemy.Update = _build_filter(sqlalchemy.update(Event)) + try_lock_cte = ( + try_lock + .where(sqlalchemy.or_( + Event.locked_until == None, + Event.locked_until < int(time_now.timestamp()) + )) + .values( + locked_until = int(locked_until.timestamp()), + locked_by = lock_owner_token + ) + .returning(Event) + ).cte("try_lock_cte") + + check_lock = sqlalchemy.exists( + sqlalchemy.select(try_lock_cte.c.id) \ + .where(try_lock_cte.c.id == id) \ + .where(try_lock_cte.c.locked_by == lock_owner_token) + ) + + stmt = sqlalchemy.select( + sqlalchemy.case( + (check_lock, "success"), + else_="locked" + ), + Event + ) \ + .options(selectinload(Event.users)) \ + .join(check_exists_cte, check_exists_cte.c.id == Event.id) \ + .where(check_exists_cte.c.id == id) + + # execute + async with session.begin(): + results = (await session.execute(stmt)).all() + if len(results) == 0: + raise NotFoundError + else: + status = results[0][0] + event_db = results[0][1] + if status == "success": + return event_db, lock_owner_token + elif status == "locked": + raise LockedError + else: + raise RuntimeError("unexpected lock status") + + +async def read_event_many( + session:AsyncSession, + type:Literal["ctftime", "custom"], + archived:Optional[bool]=None, + limit:Optional[int]=None, + # ctftime events finish_after:Optional[int]=None, + finish_before:Optional[int]=None, + # ctftime events (finish_before mode) and custom events + before_id:Optional[int]=None, ) -> List[Event]: """ Read Events. + There are two types of Events: + - ``type=ctftime`` + - finish_after mode + - Search Events which are finish after ``finish_after`` + - finish_before mode + - Search Events (1) which are finish before ``finish_before`` (2) with id smaller than ``before_id`` + - ``finish_before`` - ``finish`` of the last event in previous page + - ``before_id`` - ``id`` of the last event in previous page + - ``finish_before=None`` and ``before_id=None`` for "first page" + - ``limit`` is required + - ``type=custom`` + - Search Events with id smaller than ``before_id`` + - ``before_id=None`` for "first page" + - ``limit`` is required + :param session: - :param lock_owner_token: - :param type: Search "ctftime", "custom" Events, or ``None`` to search both types of Events. + :param type: Search ``ctftime`` or ``custom`` Events. :param archived: Search archived, non-archived Events, or ``None`` to search both types of Events. - :param id: - :param channel_id: Discord channel ID. - :param event_id: (CTFTime Event only) CTFTime Event id. - :param finish_after: (CTFTime Event only) search CTFTime Events which finish after ``finish_after``. + :param limit: + :param finish_after: + :param finish_before: + :param before_id: :return List[Event]: A list of Events. - :raise ValueError: Invalid arguments. - :raise RuntimeError: + :raise ValueError: :raise (Exception from sqlalchemy): """ # stmt @@ -316,56 +526,64 @@ async def read_event( .options(selectinload(Event.users)) # arguments - if type is not None: - if type == "ctftime": - if event_id is not None: - stmt = stmt.where(Event.event_id == event_id) - else: - stmt = stmt.where(Event.event_id != None) - - if finish_after is not None: - stmt = stmt.where(Event.finish >= finish_after) + if type == "ctftime": + stmt = stmt.where(Event.event_id != None) \ + .order_by(sqlalchemy.desc(Event.finish), sqlalchemy.desc(Event.id)) + + if finish_after is not None: + # finish_after mode + if (finish_before is not None) or (limit is not None) or (before_id is not None): + raise ValueError("finish_before, limit and before_id are not available for CTFTime Events in finish_after mode") - # sorting - stmt = stmt.order_by(sqlalchemy.asc(Event.finish)) - elif type == "custom": - stmt = stmt.where(Event.event_id == None) + stmt = stmt.where(Event.finish >= finish_after) else: - raise ValueError(f"type should be \"ctftime\", \"custom\" or None") + # finish_before mode + + # limit + if (limit is None) or (limit <= 0): + raise ValueError("limit is required and must be greater than 0 for CTFTime Events in finish_before mode") + + stmt = stmt.limit(limit) + + # finish_before & before_id + if (finish_before is not None) and (before_id is not None): + stmt = stmt.where(sqlalchemy.or_( + Event.finish < finish_before, + sqlalchemy.and_( + Event.finish == finish_before, + Event.id < before_id + ) + )) + else: + if (finish_before is None) and (before_id is None): + # first page + pass + else: + raise ValueError("invalid finish_before and before_id for CTFTime Events in finish_before mode") + elif type == "custom": + if (finish_after is not None) or (finish_before is not None): + raise ValueError("finish_after and finish_before are not available for custom Events") + + if (limit is None) or (limit <= 0): + raise ValueError("limit is required and must be greater than 0 for custom Events") + + stmt = stmt.where(Event.event_id == None) \ + .order_by(sqlalchemy.desc(Event.id)) \ + .limit(limit) + + if before_id is not None: + stmt = stmt.where(Event.id < before_id) + else: + raise ValueError("invalid type") if archived is not None: stmt = stmt.where(Event.archived == archived) - - if id is not None: - stmt = stmt.where(Event.id == id) - - if channel_id is not None: - stmt = stmt.where(Event.channel_id == channel_id) - + # execute try: - results = (await session.execute(stmt)).scalars().all() + return (await session.execute(stmt)).scalars().all() except Exception: raise - - # check lock - if lock_owner_token is not None: - # only effective when conditions include unique columes. - if (id is not None or \ - channel_id is not None or \ - event_id is not None) and \ - len(results) == 1: - result = results[0] - - time_now = datetime.now(timezone.utc) - - if result.locked_by is None or \ - result.locked_by != lock_owner_token or \ - result.locked_until is None or \ - result.locked_until < int(time_now.timestamp()): - raise RuntimeError("Invalid lock") - - return results async def read_ctfime_events_need_archive(session:AsyncSession, finish_before:int) -> List[Event]: @@ -456,6 +674,7 @@ async def update_event( try: result = (await session.execute(stmt)).scalar_one() await session.flush() + await session.refresh(result) return result except Exception: - raise \ No newline at end of file + raise diff --git a/src/crud/user.py b/src/crud/user.py index 107e524..aa26a59 100644 --- a/src/crud/user.py +++ b/src/crud/user.py @@ -29,6 +29,7 @@ async def create_user(session:AsyncSession, discord_id:int) -> User: try: result = (await session.execute(stmt)).scalar_one() await session.flush() + await session.refresh(result) return result except Exception: raise @@ -103,6 +104,7 @@ async def update_user( try: result = (await session.execute(stmt)).scalar_one() await session.flush() + await session.refresh(result) return result except Exception: raise \ No newline at end of file diff --git a/src/router/__init__.py b/src/router/__init__.py index 7f570c6..a9c585f 100644 --- a/src/router/__init__.py +++ b/src/router/__init__.py @@ -2,3 +2,4 @@ from .user import router as user_router from .ctf import router as ctf_router from .config import router as config_router +from .guild import router as guild_router \ No newline at end of file diff --git a/src/router/config.py b/src/router/config.py index a5a5e8c..54a5d01 100644 --- a/src/router/config.py +++ b/src/router/config.py @@ -1,12 +1,11 @@ from typing import Optional import logging -from fastapi import APIRouter, HTTPException, Request, Depends +from fastapi import APIRouter, HTTPException, Depends import discord from src.backend import security from src.backend import config as config_backend -from src.bot import get_bot from src import schema # logger @@ -22,9 +21,8 @@ async def read_config( key:Optional[str]=None, member:discord.Member=Depends(security.fastapi_check_administrator) ) -> schema.ConfigResponse: - bot = await get_bot() try: - config_info = await config_backend.read_config(bot, key) + config_info = await config_backend.read_config(key) except HTTPException: raise except Exception as e: @@ -42,9 +40,8 @@ async def update_config( member:discord.Member=Depends(security.fastapi_check_administrator), ) -> schema.General: # update config - bot = await get_bot() try: - await config_backend.update_config(bot, (key, data.value)) + await config_backend.update_config((key, data.value)) except HTTPException: raise except Exception as e: diff --git a/src/router/ctf.py b/src/router/ctf.py index 448d0bc..482f293 100644 --- a/src/router/ctf.py +++ b/src/router/ctf.py @@ -1,7 +1,7 @@ from typing import Optional, List, Literal import logging -from fastapi import APIRouter, Depends, HTTPException +from fastapi import APIRouter, Depends, HTTPException, Query from sqlalchemy.ext.asyncio import AsyncSession import discord @@ -9,7 +9,9 @@ from src.backend import security from src.backend import channel_op from src.backend import event as event_backend +from src.bot import get_guild from src import schema +from src import crud # logger logger = logging.getLogger("uvicorn") @@ -35,26 +37,65 @@ async def create_custom_event( # read -@router.get("/") -async def read_event_all( - type:Literal["ctftime", "custom"], +@router.get("/ctftime") +async def read_all_ctftime_event( archived:Optional[bool]=None, + limit:int=Query(gt=0, le=20), + finish_before:Optional[int]=Query(ge=0, default=None), + before_id:Optional[int]=Query(ge=0, default=None), session:AsyncSession=Depends(fastapi_get_db), - member:discord.Member=Depends(security.fastapi_check_user) + member:discord.Member=Depends(security.fastapi_check_user), ) -> List[schema.Event]: + # argument check + first_page = (finish_before is None) and (before_id is None) + n_page = (finish_before is not None) and (before_id is not None) + if first_page == False and n_page == False: + raise HTTPException(400, "invalid finish_before and before_id") + + # get events from database try: - events = await event_backend.get_event( + events_db = await crud.read_event_many( session=session, - type=type, - archived=archived + type="ctftime", + archived=archived, + limit=limit, + finish_after=None, + finish_before=finish_before, + before_id=before_id + ) + except Exception as e: + logger.error(f"fail to read Events from database: {str(e)}") + raise HTTPException(500, "fail to read Events from database") + + # format and return + return (await event_backend.format_event(get_guild(), events_db)) + + +@router.get("/custom") +async def read_all_custom_event( + archived:Optional[bool]=None, + limit:int=Query(gt=0, le=20), + before_id:Optional[int]=Query(ge=0, default=None), + session:AsyncSession=Depends(fastapi_get_db), + member:discord.Member=Depends(security.fastapi_check_user), +) -> List[schema.Event]: + # get events from database + try: + events_db = await crud.read_event_many( + session=session, + type="custom", + archived=archived, + limit=limit, + finish_after=None, + finish_before=None, + before_id=before_id ) - except HTTPException: - raise except Exception as e: - logger.error(f"fail to read Events: {str(e)}") - raise HTTPException(500, "fail to read Events") + logger.error(f"fail to read Events from database: {str(e)}") + raise HTTPException(500, "fail to read Events from database") - return events + # format and return + return (await event_backend.format_event(get_guild(), events_db)) @router.get("/{event_db_id}") @@ -62,16 +103,22 @@ async def read_event( event_db_id:int, session:AsyncSession=Depends(fastapi_get_db), member:discord.Member=Depends(security.fastapi_check_user) -) -> List[schema.Event]: +) -> schema.Event: + # get event from database try: - events = await event_backend.get_event(session=session, id=event_db_id) - except HTTPException: - raise + event_db, _ = await crud.read_event_one( + session, + lock=False, + id=event_db_id + ) + except crud.NotFoundError: + raise HTTPException(404, f"Event (id={event_db_id}) not found") except Exception as e: - logger.error(f"fail to read Events: {str(e)}") - raise HTTPException(500, "fail to read Events") + logger.error(f"fail to read Event (id={event_db_id}) from database: {str(e)}") + raise HTTPException(500, f"fail to read Event (id={event_db_id}) from database") - return events + # format and return + return (await event_backend.format_event(get_guild(), [event_db]))[0] # update - join diff --git a/src/router/guild.py b/src/router/guild.py new file mode 100644 index 0000000..c3912f6 --- /dev/null +++ b/src/router/guild.py @@ -0,0 +1,59 @@ +from typing import List +import logging + +from fastapi import APIRouter, Depends +import discord + +from src.backend.security import fastapi_check_user +from src.bot import get_guild +from src import schema + +# logger +logger = logging.getLogger("uvicorn") + +# router +router = APIRouter(prefix="/guild", tags=["Guild"]) + +@router.get("/text_channels") +async def guild_text_channels( + member:discord.Member=Depends(fastapi_check_user), + guild:discord.Guild=Depends(get_guild) +) -> List[schema.DiscordTextChannel]: + results = [] + for c in guild.text_channels: + if c.permissions_for(member).view_channel == True: + results.append(schema.DiscordTextChannel( + id=c.id, + jump_url=c.jump_url, + name=c.name + )) + return results + + +@router.get("/categories") +async def guild_categories( + member:discord.Member=Depends(fastapi_check_user), + guild:discord.Guild=Depends(get_guild), +) -> List[schema.DiscordCategoryChannel]: + results = [] + for c in guild.categories: + if c.permissions_for(member).view_channel == True: + results.append(schema.DiscordCategoryChannel( + id=c.id, + jump_url=c.jump_url, + name=c.name + )) + return results + + +@router.get("/roles") +async def guild_roles( + member:discord.Member=Depends(fastapi_check_user), + guild:discord.Guild=Depends(get_guild), +) -> List[schema.DiscordRole]: + return [ + schema.DiscordRole( + id=r.id, + name=r.name + ) for r in guild.roles + ] diff --git a/src/router/user.py b/src/router/user.py index 5963d3b..03c7e22 100644 --- a/src/router/user.py +++ b/src/router/user.py @@ -22,11 +22,17 @@ # delete - (nope) @router.get("/") +async def read_all_user( + session:AsyncSession=Depends(fastapi_get_db), + member:discord.Member=Depends(fastapi_check_user) +) -> List[schema.User]: + return (await user.get_user(session)) + + @router.get("/{discord_id}") async def read_user( - discord_id:Optional[int]=None, + discord_id:int, session:AsyncSession=Depends(fastapi_get_db), member:discord.Member=Depends(fastapi_check_user) -) -> List[schema.User]: - users = await user.get_user(session, discord_id) - return users +) -> schema.User: + return (await user.get_user(session, discord_id))[0] diff --git a/src/schema/__init__.py b/src/schema/__init__.py index 97d6594..d61a4ae 100644 --- a/src/schema/__init__.py +++ b/src/schema/__init__.py @@ -1,7 +1,10 @@ from .general import General from .config import Config, ConfigResponse, UpdateConfig -from .user import UserSimple, User, DiscordUser, UpdateUser -from .event import DiscordChannel, EventSimple, Event, CreateCustomEvent, RelinkEvent +from .user import UserRole, UserSimple, User, DiscordUser, UpdateUser +from .event import EventSimple, Event, CreateCustomEvent, RelinkEvent +from .guild import DiscordTextChannel, DiscordCategoryChannel, DiscordRole User.model_rebuild() + +EventSimple.model_rebuild() Event.model_rebuild() diff --git a/src/schema/event.py b/src/schema/event.py index ec534d4..b627041 100644 --- a/src/schema/event.py +++ b/src/schema/event.py @@ -14,12 +14,6 @@ class RelinkEvent(BaseModel): # response schema -class DiscordChannel(BaseModel): - id:int - jump_url:str - name:str - - class EventSimple(BaseModel): model_config = ConfigDict(from_attributes=True) @@ -32,7 +26,7 @@ class EventSimple(BaseModel): finish:Optional[int]=None channel_id:Optional[int]=None - channel:Optional[DiscordChannel]=None + channel:Optional["DiscordTextChannel"]=None scheduled_event_id:Optional[int]=None # extra attrbutes diff --git a/src/schema/guild.py b/src/schema/guild.py new file mode 100644 index 0000000..85b72dd --- /dev/null +++ b/src/schema/guild.py @@ -0,0 +1,18 @@ +from pydantic import BaseModel + +# response schema +class DiscordTextChannel(BaseModel): + id:int + jump_url:str + name:str + + +class DiscordCategoryChannel(BaseModel): + id:int + jump_url:str + name:str + + +class DiscordRole(BaseModel): + id:int + name:str \ No newline at end of file diff --git a/src/schema/user.py b/src/schema/user.py index eb3b6c8..6cf3b82 100644 --- a/src/schema/user.py +++ b/src/schema/user.py @@ -1,5 +1,6 @@ from __future__ import annotations from typing import Optional, List +from enum import Enum from pydantic import BaseModel, ConfigDict @@ -13,6 +14,12 @@ class UpdateUser(BaseModel): # get +class UserRole(Enum): + administrator="Administrator" + pm="pm" + member="member" + + class DiscordUser(BaseModel): display_name:str id:int @@ -23,6 +30,7 @@ class UserSimple(BaseModel): model_config = ConfigDict(from_attributes=True) discord_id:int + user_role:List[UserRole] status:model.Status skills:List[model.Skills] diff --git a/src/utils/notification.py b/src/utils/notification.py index 7f7725d..f8bef0e 100644 --- a/src/utils/notification.py +++ b/src/utils/notification.py @@ -3,7 +3,7 @@ import discord -from src.bot import get_bot +from src.bot import get_guild from src.config import settings async def send_notification( @@ -22,9 +22,9 @@ async def send_notification( :raise RuntimeError: """ - bot = await get_bot() - guild = bot.get_guild(settings.GUILD_ID) - if guild is None: + try: + guild = get_guild() + except Exception: raise RuntimeError(f"Guild (id={settings.GUILD_ID}) not found") # args