diff --git a/pyproject.toml b/pyproject.toml index d95208c..465104e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,6 +6,7 @@ build-backend = "setuptools.build_meta" name = "fregbot" dynamic = ["version"] dependencies = [ + "click>=8.3.1", "discord.py", "pydantic", ] diff --git a/src/spielverabredungen.py b/src/spielverabredungen.py index 39ead24..f4e3a2e 100644 --- a/src/spielverabredungen.py +++ b/src/spielverabredungen.py @@ -1,15 +1,20 @@ import asyncio import datetime +import itertools import locale import logging +from collections import defaultdict from datetime import timedelta +from functools import cached_property from typing import Iterable import click import discord -from discord import Client, ForumChannel +from discord import Client, ForumChannel, Guild, Thread, ScheduledEvent from pydantic import BaseModel, PositiveInt +AVAILABLE_TABLES = 11 + class SpielverabredungenConfig(BaseModel): guild_id: int @@ -18,14 +23,11 @@ class SpielverabredungenConfig(BaseModel): end: datetime.datetime go_back_threads_days: PositiveInt - @property - def dates(self) -> Iterable[datetime.date]: + @cached_property + def dates(self) -> tuple[datetime.date, ...]: start, end = self.start.date(), self.end.date() assert start <= end, f'Start date {self.start} must be smaller than end date {self.end}' - curr_date = start - while curr_date <= end: - yield curr_date - curr_date += timedelta(days=1) + return tuple(start + timedelta(days=n) for n in range((end - start).days + 1)) class SpielverabredungenClient(Client): @@ -42,10 +44,10 @@ class SpielverabredungenClient(Client): logging.info('Channel name is "%s"', channel.name) return channel - def __check_duplicate_channels(self, channel: ForumChannel, dates: Iterable[datetime.date], - thread_lookback_days: int) -> Iterable[datetime.date]: - skipped_dates = 0 - potential_threadnames = {self.__get_thread_name(date): date for date in dates} + def __resolve_threads(self, channel: ForumChannel, dates: Iterable[datetime.date], thread_lookback_days: int + ) -> dict[datetime.date, Thread | None]: + threads: dict[datetime.date, Thread | None] = {date: None for date in dates} + potential_threadnames = {self.__get_thread_name(date): date for date in threads.keys()} max_lookback_date = datetime.date.today() - datetime.timedelta(days=thread_lookback_days) for thread in channel.threads: if thread.locked: @@ -53,34 +55,87 @@ class SpielverabredungenClient(Client): if thread.created_at.date() < max_lookback_date: break if thread.name in potential_threadnames: - del potential_threadnames[thread.name] - skipped_dates += 1 - logging.info('Skipping %d days from posting, as the threads were already found', skipped_dates) - yield from potential_threadnames.values() + threads[potential_threadnames[thread.name]] = thread + logging.info('Found %d of %d threads already existing. Will update', + sum(val is not None for val in threads.values()), len(threads)) + return threads + + def __handle_longrunning_events(self, event: ScheduledEvent, date_offset: timedelta = timedelta(days=0)) -> Iterable[tuple[datetime.date, ScheduledEvent]]: + curr_date = event.start_time.date() + date_offset + while curr_date <= (event.end_time.date() + date_offset): + yield curr_date, event + curr_date += timedelta(days=1) + + def __handle_recurring_events(self, event: ScheduledEvent) -> Iterable[tuple[datetime.date, ScheduledEvent]]: + """TODO fix after https://github.com/Rapptz/discord.py/pull/9685 is merged""" + if event.name != 'Mal- und Basteltreff': + yield from self.__handle_longrunning_events(event) + return + week_counter = 0 + while (curr_date := event.start_time.date() + timedelta(weeks=week_counter)) <= self.config_object.end.date(): + if curr_date >= self.config_object.start.date(): + yield from self.__handle_longrunning_events(event, timedelta(weeks=week_counter)) + week_counter += 1 + + def __get_relevant_events(self, guild: Guild) -> dict[datetime.date, list[ScheduledEvent]]: + relevant_events = defaultdict(list) + for date, event in itertools.chain.from_iterable(map(self.__handle_recurring_events, guild.scheduled_events)): + if date > self.config_object.end.date(): + continue + relevant_events[date].append(event) + return dict(relevant_events) @staticmethod def __get_thread_name(date: datetime.date) -> str: return date.strftime('%Y-%m-%d %A') - async def __create_thread(self, channel: ForumChannel, date: datetime.date) -> bool: + def __reservation_text(self, date: datetime.date, relevant_events: list[ScheduledEvent]) -> str: + content = f"Tragt hier eure Spielverabredungen ein für: {self.__get_thread_name(date)}\n" \ + "Reservierungen in diesem Kanal haben Priorität gegenüber spontanen Spielen.\n" \ + f"Es gibt insgesamt **{AVAILABLE_TABLES}** Tische.\n" + if not relevant_events: + return content + content += "\n\nBeachtet bitte ggf. den Platz für die folgenden Termine:\n" + for relevant_event in sorted(relevant_events, key=lambda event: event.start_time): + content += f"[{relevant_event.name}]({relevant_event.url})\n" + return content + + async def __create_thread(self, channel: ForumChannel, date: datetime.date, content: str) -> bool: name = self.__get_thread_name(date) - thread = await channel.create_thread(name=name, - content="Tragt hier eure Spielverabredungen ein für oben genanntes Datum.", - reason='Spielverabredungen - Erstellt von Bot') + thread = await channel.create_thread(name=name, content=content, reason='Spielverabredungen - Erstellt von Bot', + silent=True) return thread is not None + async def __update_thread(self, thread: Thread, content: str) -> bool: + async for message in thread.history(oldest_first=True): + if message.author != self.user: + logging.info("Author of thread %s is not the bot user. Cannot update", thread.jump_url) + return False + if message.content != content: + await message.edit(content=content) + break + return True + + + async def on_ready(self): - logging.info(f'Logging in as {self.user}') + logging.info('Logging in as %s', self.user) try: channel = self.__get_channel() logging.info('Channel obtained') - dates = self.__check_duplicate_channels(channel, self.config_object.dates, - self.config_object.go_back_threads_days) - for index, date in enumerate(dates, start=1): - if index % 10 == 0: - logging.info('Attempting Thread creation number %d for %s', index, date) - if not await self.__create_thread(channel, date): - logging.error('Could not create thread for %s', date) + relevant_events = self.__get_relevant_events(channel.guild) + threads = self.__resolve_threads(channel, self.config_object.dates, + self.config_object.go_back_threads_days) + for index, (date, thread) in enumerate(threads.items(), start=1): + if index % 5 == 0: + logging.info('Handling thread %d of %d for %s', index, len(threads), date) + content = self.__reservation_text(date, relevant_events.get(date, [])) + if thread is None: + if not await self.__create_thread(channel, date, content): + logging.error('Could not create thread for %s', date) + else: + if not await self.__update_thread(thread, content): + logging.error("Could not update thread for %s", date) await asyncio.sleep(1) except Exception as e: logging.error(f'Error: {e}') @@ -105,6 +160,7 @@ def main(guild: int, channel: int, start: datetime.datetime, end: datetime.datet thread_lookback: int, debug: bool) -> None: intents = discord.Intents.default() intents.messages = True + intents.guild_scheduled_events = True log_handler = logging.FileHandler(filename='spielabredungen.log', encoding='utf-8', mode='w') if debug: