diff --git a/database.py b/database.py new file mode 100644 index 0000000..731366b --- /dev/null +++ b/database.py @@ -0,0 +1,199 @@ +# database.py +import sqlite3 +from datetime import datetime +from pathlib import Path +from typing import Any, Dict, Tuple + +DB_FILE = "data/database.db" + + +def ensure_db() -> None: + """Initialize SQLite database and create tables if needed.""" + Path(DB_FILE).parent.mkdir(parents=True, exist_ok=True) + + with sqlite3.connect(DB_FILE) as conn: + cursor = conn.cursor() + cursor.execute(""" + CREATE TABLE IF NOT EXISTS guild_state ( + guild_id TEXT PRIMARY KEY, + guild_name TEXT NOT NULL DEFAULT '', + cooldown_until TEXT NOT NULL DEFAULT '1970-01-01T00:00:00', + self_control REAL NOT NULL DEFAULT 1.0, + last_updated TEXT NOT NULL DEFAULT (datetime('now')) + ) + """) + cursor.execute(""" + CREATE TABLE IF NOT EXISTS active_riddles ( + channel_id TEXT PRIMARY KEY, + riddle_index INTEGER NOT NULL, + nb_clues INTEGER NOT NULL DEFAULT -1, + message_id INTEGER NOT NULL, + solver_id TEXT, + created_at TEXT NOT NULL DEFAULT (datetime('now')), + updated_at TEXT NOT NULL DEFAULT (datetime('now')) + ) + """) + conn.commit() + + +def get_connection() -> sqlite3.Connection: + """Return SQLite connection with row factory for named column access.""" + conn = sqlite3.connect(DB_FILE) + conn.row_factory = sqlite3.Row + return conn + + +def get_guild_state(guild_id: str) -> Dict[str, Any]: + """Retrieve guild state from database.""" + with get_connection() as conn: + cursor = conn.cursor() + cursor.execute( + """ + SELECT guild_id, guild_name, cooldown_until, self_control, last_updated + FROM guild_state WHERE guild_id = ? + """, + (guild_id,), + ) + row = cursor.fetchone() + + if row: + return { + "guild_id": row["guild_id"], + "guild_name": row["guild_name"], + "cooldown_until": row["cooldown_until"], + "self_control": row["self_control"], + "last_updated": row["last_updated"], + } + else: + return { + "guild_id": guild_id, + "guild_name": "", + "cooldown_until": "1970-01-01T00:00:00", + "self_control": 1.0, + "last_updated": datetime.now().isoformat(), + } + + +def update_guild_state( + guild_id: str, guild_name: str, cooldown_until: str, self_control: float +) -> None: + """Update guild state in database.""" + with get_connection() as conn: + cursor = conn.cursor() + cursor.execute( + """ + INSERT OR REPLACE INTO guild_state (guild_id, guild_name, cooldown_until, self_control, last_updated) + VALUES (?, ?, ?, ?, ?) + """, + ( + guild_id, + guild_name, + cooldown_until, + self_control, + datetime.now().isoformat(), + ), + ) + conn.commit() + + +def delete_guild_state(guild_id: str) -> bool: + """Delete guild state from database.""" + with get_connection() as conn: + cursor = conn.cursor() + cursor.execute("DELETE FROM guild_state WHERE guild_id = ?", (guild_id,)) + conn.commit() + return cursor.rowcount > 0 + + +def get_all_guild_states() -> Dict[str, Dict[str, Any]]: + """Retrieve all guild states (for debug purposes).""" + with get_connection() as conn: + cursor = conn.cursor() + cursor.execute( + "SELECT guild_id, guild_name, cooldown_until, self_control, last_updated FROM guild_state" + ) + return { + row["guild_id"]: { + "guild_name": row["guild_name"], + "cooldown_until": row["cooldown_until"], + "self_control": row["self_control"], + "last_updated": row["last_updated"], + } + for row in cursor.fetchall() + } + + +def get_active_riddle(channel_id: str): + """Récupérer une énigme active pour un canal spécifique.""" + with get_connection() as conn: + cursor = conn.cursor() + cursor.execute( + "SELECT * FROM active_riddles WHERE channel_id = ?", (str(channel_id),) + ) + row = cursor.fetchone() + + if row: + return { + "channel_id": row["channel_id"], + "riddle_index": row["riddle_index"], + "nb_clues": row["nb_clues"], + "message_id": row["message_id"], + "solver_id": row["solver_id"], + "created_at": row["created_at"], + "updated_at": row["updated_at"], + } + return None + + +def save_active_riddle( + channel_id: str, + riddle_index: int, + nb_clues: int, + message_id: int, + solver_id: str = None, +): + """Sauvegarder ou mettre à jour une énigme active.""" + with get_connection() as conn: + cursor = conn.cursor() + # Si solver_id existe, on considère que c'est résolu et on garde l'entrée pour historique + # Sinon on met à jour l'énigme en cours + + cursor.execute( + """ + INSERT OR REPLACE INTO active_riddles + (channel_id, riddle_index, nb_clues, message_id, solver_id, updated_at) + VALUES (?, ?, ?, ?, ?, datetime('now')) + """, + (str(channel_id), riddle_index, nb_clues, message_id, solver_id), + ) + conn.commit() + + +def remove_active_riddle(channel_id: str) -> bool: + """Supprimer une énigme active du suivi.""" + with get_connection() as conn: + cursor = conn.cursor() + cursor.execute( + "DELETE FROM active_riddles WHERE channel_id = ?", (str(channel_id),) + ) + conn.commit() + return cursor.rowcount > 0 + + +def get_all_active_riddles(): + """Récupérer toutes les énigmes actives (pour debug).""" + with get_connection() as conn: + cursor = conn.cursor() + cursor.execute("SELECT * FROM active_riddles") + return [ + { + "channel_id": row["channel_id"], + "riddle_index": row["riddle_index"], + "nb_clues": row["nb_clues"], + "message_id": row["message_id"], + "solver_id": row["solver_id"], + "created_at": row["created_at"], + "updated_at": row["updated_at"], + } + for row in cursor.fetchall() + ] diff --git a/fouras.py b/fouras.py index 916a67b..505b68b 100644 --- a/fouras.py +++ b/fouras.py @@ -1,31 +1,32 @@ -# fouras.py +# Remplacez les imports globaux et initialisations +import json import random import re -import json +from typing import Any, Dict + import discord from unidecode import unidecode -from pathlib import Path -from typing import Dict, Any, Tuple -API_URL = "".join([ - "https://discord.com/api/oauth2/authorize?", - "client_id=1110208055171367014&permissions=274877975552&scope=bot", -]) +import database + +# Supprimez ces lignes globales : +# ongoing_riddles = {} +# riddles = [] +# answers = [] + +# Gardez seulement le chargement +riddles = [] +answers = [] + +API_URL = "".join( + [ + "https://discord.com/api/oauth2/authorize?", + "client_id=1110208055171367014&permissions=274877975552&scope=bot", + ] +) MAINTAINER_ID = 151626081458192384 -BUG_REPORT = """ -BUG REPORT from {user} (`{user_id}`) in channel {channel} (`{channel_id}`) : - -Message : -> {message} - -State : -```json\n{state}``` -History : -```json\n{history}``` -""" - ABOUT = """ Ce bot a été développé par {user} Code Source : https://git.epicsparrow.com/Anselme/perefouras @@ -42,16 +43,6 @@ Numéro d'énigme invalide, merci de saisir un numéro entre 1 et {len} RIDDLES_FILE = "resources/riddles.txt" ANSWERS_FILE = "resources/answers.txt" -SAVE_FILE = "data/fouras_riddles.json" - -riddles = [] -answers = [] -ongoing_riddles = {} - - -def _ensure_data_dir() -> None: - """Ensure data directory exists.""" - Path(SAVE_FILE).parent.mkdir(parents=True, exist_ok=True) def load_riddles() -> list: @@ -66,42 +57,8 @@ def load_riddles() -> list: print(f"Loaded {len(riddles)} riddles") -def load_ongoing_riddles(client) -> Dict[str, Dict[str, Any]]: - """Load ongoing riddles from save file.""" - try: - with open(SAVE_FILE, "r", encoding="utf-8") as f: - config = json.load(f) - - ongoing_riddles = {} - for channel_id, channel_info in config.items(): - channel = client.fetch_channel(int(channel_id)) - channel_info["message"] = channel.fetch_message(channel_info["message"]) - ongoing_riddles[channel] = channel_info - - return ongoing_riddles - except FileNotFoundError: - return {} - except json.JSONDecodeError: - return {} - - -def save_ongoing_riddles(ongoing_riddles: Dict) -> str: - """Save ongoing riddles state to file.""" - _ensure_data_dir() - dump = {} - for key, value in ongoing_riddles.items(): - dump_channel = dict(value) - dump_channel["message"] = dump_channel["message"].id - dump[key.id] = dump_channel - - with open(SAVE_FILE, "w", encoding="utf-8") as f: - json.dump(dump, f, ensure_ascii=False, indent=2) - - return f'Saved fouras riddles state in file "{SAVE_FILE}"' - - -def new_riddle(riddles: list, answers: list, index: int) -> Dict[str, Any]: - """Create a new riddle state.""" +def new_riddle_state(index: int) -> Dict[str, Any]: + """Créer un nouvel état d'énigme sans sauvegarde DB.""" return { "index": index, "nbClues": -1, @@ -110,10 +67,13 @@ def new_riddle(riddles: list, answers: list, index: int) -> Dict[str, Any]: } -def finish_riddle(ongoing_riddles: Dict, channel) -> None: - """Remove completed riddle from tracking.""" - if channel in ongoing_riddles: - del ongoing_riddles[channel] +async def finish_riddle(channel, client) -> None: + """Supprimer l'énigme active de la DB et supprimer le suivi en mémoire.""" + removed = database.remove_active_riddle(str(channel.id)) + + # Optionnel : si vous gardez un cache en mémoire pour performance, supprimez-le ici + # if hasattr(client, '_cache_ongoing_riddles'): + # client._cache_ongoing_riddles.pop(channel.id, None) def clue_string(answer: str, nb_clues: int) -> str: @@ -121,44 +81,53 @@ def clue_string(answer: str, nb_clues: int) -> str: final_string = "_" for _ in range(len(answer) - 1): final_string += " _" - + random.seed(hash(answer)) nb_revealed = 0 - + for _ in range(nb_clues): idx = random.randint(0, len(answer) - 1) while final_string[idx * 2] != "_": idx = random.randint(0, len(answer) - 1) - + nb_revealed += 1 - final_string = final_string[:idx * 2] + answer[idx] + final_string[idx * 2 + 1:] - + final_string = ( + final_string[: idx * 2] + answer[idx] + final_string[idx * 2 + 1 :] + ) + if nb_revealed == len(answer): return final_string - + return final_string -def format_riddle_message(current_riddle: Dict[str, Any]) -> str: +def format_riddle_message( + current_riddle: Dict[str, Any], solved: bool = False, solver_mention: str = "" +) -> str: """Format riddle message for display.""" - nb_clues = current_riddle["nbClues"] - answer = current_riddle["answer"] - + nb_clues = current_riddle.get("nbClues", -1) + answer = current_riddle.get("answer", "") + + # Si l'énigme est résolue, on récupère le solver depuis l'ID + solver = None + if solved and solver_mention: + solver = type("obj", (object,), {"mention": solver_mention})() + formatted_riddle = "> " + current_riddle["riddle"].replace("\n", "\n> ") formatted_riddle = formatted_riddle.replace("\r", "") - + clue = "" if nb_clues > -1: if nb_clues >= len(answer): clue = "\nNon trouvée, la solution était : `{0}`".format(answer) else: clue = "\nIndice : `{0}`".format(clue_string(answer, nb_clues)) - - if "solver" in current_riddle: + + if solved and solver: clue = clue + "\n{0} a trouvé la solution, qui était : `{1}`".format( - current_riddle["solver"].mention, answer + solver.mention, answer ) - + if clue: return "Énigme {0}:\n{1}\n> Qui suis-je ?\n{2}".format( current_riddle["index"] + 1, formatted_riddle, clue @@ -168,64 +137,79 @@ def format_riddle_message(current_riddle: Dict[str, Any]) -> str: current_riddle["index"] + 1, formatted_riddle ) + async def get_channel_name(channel, client) -> str: if isinstance(channel, discord.DMChannel): dm_channel = await client.fetch_channel(channel.id) return "[DM={0}]".format(dm_channel.recipient.name) else: - return "[Server={0}] => [Channel={1}]".format( - channel.guild.name, channel.name - ) + return "[Server={0}] => [Channel={1}]".format(channel.guild.name, channel.name) -async def handle_debug_commands(message, client, ongoing_riddles: Dict) -> bool: + +async def handle_debug_commands(message, client) -> bool: """Handle debug commands (debug, save, load, broadcast). Returns True if handled.""" message_content = message.content.lower() - - if message_content == "save fouras": - if message.author.id == MAINTAINER_ID: - status = save_ongoing_riddles(ongoing_riddles) - json_str = "```json\n{0}```".format( - json.dumps(ongoing_riddles, ensure_ascii=False, indent=2) - ) - await message.author.send(status) - await message.author.send(json_str) - return True - + if message_content == "load fouras": if message.author.id == MAINTAINER_ID: - # Reload would require passing riddles/answers back - await message.author.send("Reloaded riddles from files") + load_riddles() + await message.author.send("Réponses et énigmes rechargées") return True - + if message_content == "debug fouras": if message.author.id == MAINTAINER_ID: + # Récupération depuis la DB au lieu d'une variable globale + db_riddles = database.get_all_active_riddles() + dump = {} - for key, value in ongoing_riddles.items(): - dump_channel = dict(value) - dump_channel.pop("message", None) - dump_channel.pop("answer", None) - channel_name = await get_channel_name(key, client) - dump[channel_name] = dump_channel + for riddle_data in db_riddles: + try: + channel = await client.fetch_channel(int(riddle_data["channel_id"])) + channel_name = await get_channel_name(channel, client) + + # Reconstruction partielle pour l'affichage + riddle_info = { + "riddle_index": riddle_data["riddle_index"], + "nb_clues": riddle_data["nb_clues"], + "message_id": riddle_data["message_id"], + "solver_id": riddle_data["solver_id"], + "created_at": riddle_data["created_at"], + "updated_at": riddle_data["updated_at"], + } + + # On ajoute le texte de l'énigme si disponible + if 0 <= riddle_data["riddle_index"] < len(riddles): + riddle_info["text"] = ( + riddles[riddle_data["riddle_index"]][:50] + "..." + ) + + dump[channel_name] = riddle_info + except Exception as e: + dump[f"Channel_{riddle_data['channel_id']}"] = {"error": str(e)} + await message.author.send( "```json\n{0}```".format(json.dumps(dump, ensure_ascii=False, indent=4)) ) return True - + broadcast_match = re.match(r"^broadcast\s+(\d+) (.*)", message.content) if broadcast_match and message.author.id == MAINTAINER_ID: index = int(broadcast_match.group(1)) broadcast_message = broadcast_match.group(2) - channel = await client.fetch_channel(index) - if channel: - await channel.send(broadcast_message) - else: - await message.channel.send(f"Invalid channel id : {index}") + try: + channel = await client.fetch_channel(index) + if channel: + await channel.send(broadcast_message) + else: + await message.channel.send(f"Invalid channel id : {index}") + except discord.errors.NotFound: + await message.channel.send(f"Channel not found : {index}") return True - + return False -async def handle_riddle_commands(message, riddles: list, answers: list, client) -> bool: +async def handle_riddle_commands(message, client) -> bool: """Handle /fouras commands. Returns True if command was processed.""" message_content = message.content.lower() @@ -236,89 +220,178 @@ async def handle_riddle_commands(message, riddles: list, answers: list, client) if random.random() <= 0.03: await message.channel.send("Non") else: - riddle_state = new_riddle(riddles, answers, index) - await message.channel.send(format_riddle_message(riddle_state)) + riddle_state = new_riddle_state(index) + + msg = await message.channel.send( + format_riddle_message({**riddle_state, "index": index}) + ) + + database.save_active_riddle( + channel_id=str(message.channel.id), + riddle_index=index, + nb_clues=-1, + message_id=msg.id, + solver_id=None, + ) else: await message.channel.send(INVALID_ID.format(len=len(riddles))) return True - + if message_content == "fouras": if random.random() <= 0.03: await message.channel.send("Non") elif len(riddles) > 0: index = random.randint(0, len(riddles) - 1) - riddle_state = new_riddle(riddles, answers, index) - await message.channel.send(format_riddle_message(riddle_state)) + riddle_state = new_riddle_state(index) + + msg = await message.channel.send( + format_riddle_message({**riddle_state, "index": index}) + ) + + database.save_active_riddle( + channel_id=str(message.channel.id), + riddle_index=index, + nb_clues=-1, + message_id=msg.id, + solver_id=None, + ) else: print(f'riddles : "{len(riddles)}"') return True - + if message_content == "about fouras": author_user = await client.fetch_user(MAINTAINER_ID) await message.channel.send(ABOUT.format(user=author_user.mention, url=API_URL)) return True - + return False -async def handle_riddle_solving(message, ongoing_riddles: Dict, client) -> bool: +async def handle_riddle_solving(message, client) -> bool: """Handle riddle solving logic. Returns True if riddle was solved or modified.""" - if message.channel not in ongoing_riddles: + # Vérification via la DB au lieu de la variable globale + riddle_data = database.get_active_riddle(str(message.channel.id)) + + if not riddle_data: return False - - current_riddle = ongoing_riddles[message.channel] - - if "message" not in current_riddle: - current_riddle["message"] = message + + # Reconstruire l'état de l'énigme + if 0 <= riddle_data["riddle_index"] < len(riddles): + current_riddle = { + "index": riddle_data["riddle_index"], + "nbClues": riddle_data["nb_clues"], + "riddle": riddles[riddle_data["riddle_index"]].strip(), + "answer": answers[riddle_data["riddle_index"]], + } + else: return False - + + # Si aucun message n'est associé, on ne peut pas éditer + if not riddle_data["message_id"]: + return False + answer = current_riddle["answer"] - + # Check if message contains the answer if unidecode(answer.lower()) in unidecode(message.content.lower()): - current_riddle["solver"] = message.author + solver_id = str(message.author.id) + + # Mise à jour DB : marque comme résolue avec le solver + database.save_active_riddle( + channel_id=str(message.channel.id), + riddle_index=riddle_data["riddle_index"], + nb_clues=current_riddle["nbClues"], + message_id=riddle_data["message_id"], + solver_id=solver_id, + ) + await message.channel.send( SUCCESS.format(user=message.author.mention, answer=answer) ) - await current_riddle["message"].edit( - content=format_riddle_message(current_riddle) - ) - finish_riddle(ongoing_riddles, message.channel) + + # Récupération du message original pour édition + try: + original_msg = await message.channel.fetch_message( + riddle_data["message_id"] + ) + await original_msg.edit( + content=format_riddle_message( + current_riddle, solved=True, solver_mention=message.author.mention + ) + ) + except discord.errors.NotFound: + pass + + # Suppression du suivi actif + await finish_riddle(message.channel, client) return True - + # Repeat riddle command if message.content.lower() in ["repete", "répète", "repeat"]: - current_riddle.pop("message", None) - await message.channel.send(format_riddle_message(current_riddle)) + try: + original_msg = await message.channel.fetch_message( + riddle_data["message_id"] + ) + # msg = await message.channel.send(format_riddle_message(current_riddle)) + # database.save_active_riddle( + # channel_id=str(message.channel.id), + # riddle_index=index, + # nb_clues=-1, + # message_id=msg.id, + # solver_id=None, + # ) + except discord.errors.NotFound: + pass return True - + # Clue command if message.content.lower() in ["indice", "aide", "help", "clue"]: nb_clues = current_riddle["nbClues"] + 1 current_riddle["nbClues"] = nb_clues - + if nb_clues >= len(answer): await message.channel.send( "Perdu ! La réponse était : `{0}`".format(answer) ) - finish_riddle(ongoing_riddles, message.channel) - else: - await current_riddle["message"].edit( - content=format_riddle_message(current_riddle) + # Mark as unsolved/revealed in DB + database.save_active_riddle( + channel_id=str(message.channel.id), + riddle_index=riddle_data["riddle_index"], + nb_clues=nb_clues, + message_id=riddle_data["message_id"], + solver_id="revealed", # Marqueur spécial ) + await finish_riddle(message.channel, client) + else: + # Mise à jour du nombre d'indices en DB + database.save_active_riddle( + channel_id=str(message.channel.id), + riddle_index=riddle_data["riddle_index"], + nb_clues=nb_clues, + message_id=riddle_data["message_id"], + solver_id=riddle_data["solver_id"], + ) + + try: + original_msg = await message.channel.fetch_message( + riddle_data["message_id"] + ) + await original_msg.edit(content=format_riddle_message(current_riddle)) + except discord.errors.NotFound: + pass return True - + return False -async def handle_bug_report(message, client, ongoing_riddles: Dict) -> bool: +async def handle_bug_report(message, client) -> bool: """Handle bug report command. Returns True if bug report was sent.""" if not message.content.lower().startswith("bug"): return False - + author_user = await client.fetch_user(MAINTAINER_ID) - channel_name = await client.get_channel_name(message.channel) - + channel_name = await get_channel_name(message.channel, client) + # Load message history messages = [ { @@ -329,9 +402,11 @@ async def handle_bug_report(message, client, ongoing_riddles: Dict) -> bool: async for msg in message.channel.history(limit=10) ] messages_json = json.dumps(messages, ensure_ascii=False) - - state_json = json.dumps(ongoing_riddles, ensure_ascii=False, indent=2) - + + # Récupération de l'état DB au lieu de la variable globale + active_riddles = database.get_all_active_riddles() + state_json = json.dumps(active_riddles, ensure_ascii=False, indent=2) + await author_user.send( BUG_REPORT.format( user=message.author.mention, @@ -350,26 +425,23 @@ async def handle_bug_report(message, client, ongoing_riddles: Dict) -> bool: async def handle_message(message, client) -> bool: - global riddles, answers, ongoing_riddles - """Main entry point for message handling.""" - if message.author == client.user: - return False - + database.ensure_db() + # Handle debug/admin commands first - if await handle_debug_commands(message, client, ongoing_riddles): + if await handle_debug_commands(message, client): return True - + # Handle riddle commands - if await handle_riddle_commands(message, riddles, answers, client): + if await handle_riddle_commands(message, client): return True - + # Handle bug reports - if await handle_bug_report(message, client, ongoing_riddles): + if await handle_bug_report(message, client): return True - + # Handle riddle solving - if await handle_riddle_solving(message, ongoing_riddles, client): + if await handle_riddle_solving(message, client): return True - - return False \ No newline at end of file + + return False diff --git a/main.py b/main.py index f3f8b9d..1fa5446 100644 --- a/main.py +++ b/main.py @@ -1,8 +1,9 @@ # main.py +import os + import discord from discord import app_commands from dotenv import load_dotenv -import os # Load environment variables load_dotenv() @@ -24,43 +25,43 @@ async def on_ready(): global client, tree """Initialize bot state and sync commands.""" # Import initialization functions - from fouras import load_riddles, _ensure_data_dir - from rhymes import load_rhymes, _ensure_db, _ensure_log_file - + from fouras import load_riddles + from rhymes import _ensure_log_file, load_rhymes + # Ensure directories exist - _ensure_data_dir() - _ensure_db() _ensure_log_file() - + # Load riddles and answers load_riddles() - + # Load rhymes success, msg = load_rhymes() print(msg) - + # Sync commands async for guild in client.fetch_guilds(): tree.copy_global_to(guild=guild) await tree.sync(guild=guild) - + print(f"Logged in as {client.user} on {len(client.guilds)} servers!") @client.event async def on_message(message): - global client, ongoing_riddles + global client """Handle incoming messages.""" # Ignore bot's own messages if message.author == client.user: return - + # Handle fouras module from fouras import handle_message as handle_fouras - await handle_fouras(message, ongoing_riddles, client) - + + await handle_fouras(message, client) + # Handle rhymes module from rhymes import handle_message as handle_rhymes + await handle_rhymes(message, client) @@ -70,4 +71,4 @@ if __name__ == "__main__": if not token: print("Error: DISCORD_TOKEN not found in environment variables") exit(1) - client.run(token) \ No newline at end of file + client.run(token) diff --git a/requirements.txt b/requirements.txt index 211581e..7f06eea 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,4 +3,4 @@ discord-py python-dotenv unidecode ruff -httpx \ No newline at end of file +httpx diff --git a/rhymes.py b/rhymes.py index b49bfd9..dbbf5a5 100644 --- a/rhymes.py +++ b/rhymes.py @@ -1,137 +1,44 @@ # rhymes.py -import random import json -import sqlite3 -from pathlib import Path +import random from datetime import datetime -from typing import Dict, Any, Tuple +from pathlib import Path +from typing import Any, Dict, Tuple + +import database RHYMES_FILE = "resources/rhymes.json" -DB_FILE = "data/poilau_state.db" RHYME_LOG_FILE = "data/rhyme_log.csv" loaded_rhymes = {} -def _ensure_db() -> None: - """Initialize SQLite database and create tables if needed.""" - Path(DB_FILE).parent.mkdir(parents=True, exist_ok=True) - - with sqlite3.connect(DB_FILE) as conn: - cursor = conn.cursor() - cursor.execute(""" - CREATE TABLE IF NOT EXISTS guild_state ( - guild_id TEXT PRIMARY KEY, - guild_name TEXT NOT NULL DEFAULT '', - cooldown_until TEXT NOT NULL DEFAULT '1970-01-01T00:00:00', - self_control REAL NOT NULL DEFAULT 1.0, - last_updated TEXT NOT NULL DEFAULT (datetime('now')) - ) - """) - conn.commit() - def _ensure_log_file() -> None: """Create CSV log file if it doesn't exist.""" Path(RHYME_LOG_FILE).parent.mkdir(parents=True, exist_ok=True) - + if not Path(RHYME_LOG_FILE).exists(): with open(RHYME_LOG_FILE, "w", encoding="utf-8") as f: f.write("timestamp,last_word,rhyme_triggered\n") -def _get_connection() -> sqlite3.Connection: - """Return SQLite connection with row factory for named column access.""" - conn = sqlite3.connect(DB_FILE) - conn.row_factory = sqlite3.Row - return conn - - def load_rhymes() -> Tuple[bool, str]: global loaded_rhymes """Load rhymes from JSON file. Returns (success, message).""" try: with open(RHYMES_FILE, "r", encoding="utf-8") as f: loaded_rhymes = json.load(f) - return True, f"Loaded rhymes file \"{RHYMES_FILE}\"" + return True, f'Loaded rhymes file "{RHYMES_FILE}"' except FileNotFoundError: - return False, f"No rhymes file found at \"{RHYMES_FILE}\"" + return False, f'No rhymes file found at "{RHYMES_FILE}"' except json.JSONDecodeError as e: - return False, f"Invalid JSON in \"{RHYMES_FILE}\": {e}" - - -def get_guild_state(guild_id: str) -> Dict[str, Any]: - """Retrieve guild state from database.""" - with _get_connection() as conn: - cursor = conn.cursor() - cursor.execute(""" - SELECT guild_id, guild_name, cooldown_until, self_control, last_updated - FROM guild_state WHERE guild_id = ? - """, (guild_id,)) - row = cursor.fetchone() - - if row: - return { - "guild_id": row["guild_id"], - "guild_name": row["guild_name"], - "cooldown_until": row["cooldown_until"], - "self_control": row["self_control"], - "last_updated": row["last_updated"] - } - else: - return { - "guild_id": guild_id, - "guild_name": "", - "cooldown_until": "1970-01-01T00:00:00", - "self_control": 1.0, - "last_updated": datetime.now().isoformat() - } - - -def update_guild_state( - guild_id: str, - guild_name: str, - cooldown_until: str, - self_control: float -) -> None: - """Update guild state in database.""" - with _get_connection() as conn: - cursor = conn.cursor() - cursor.execute(""" - INSERT OR REPLACE INTO guild_state (guild_id, guild_name, cooldown_until, self_control, last_updated) - VALUES (?, ?, ?, ?, ?) - """, (guild_id, guild_name, cooldown_until, self_control, datetime.now().isoformat())) - conn.commit() - - -def delete_guild_state(guild_id: str) -> bool: - """Delete guild state from database.""" - with _get_connection() as conn: - cursor = conn.cursor() - cursor.execute("DELETE FROM guild_state WHERE guild_id = ?", (guild_id,)) - conn.commit() - return cursor.rowcount > 0 - - -def get_all_guild_states() -> Dict[str, Dict[str, Any]]: - """Retrieve all guild states (for debug purposes).""" - with _get_connection() as conn: - cursor = conn.cursor() - cursor.execute("SELECT guild_id, guild_name, cooldown_until, self_control, last_updated FROM guild_state") - return { - row["guild_id"]: { - "guild_name": row["guild_name"], - "cooldown_until": row["cooldown_until"], - "self_control": row["self_control"], - "last_updated": row["last_updated"] - } - for row in cursor.fetchall() - } + return False, f'Invalid JSON in "{RHYMES_FILE}": {e}' def log_rhyme(last_word: str, rhyme_triggered: str) -> None: """Log rhyme trigger to CSV file.""" timestamp = datetime.now().isoformat() - + with open(RHYME_LOG_FILE, "a", encoding="utf-8") as f: safe_rhyme = rhyme_triggered.replace(",", ";") f.write(f"{timestamp},{last_word},{safe_rhyme}\n") @@ -161,17 +68,19 @@ def find_rhyme(word: str) -> str: return random.choice(rhyme["rhymes"]) return "" + async def get_guild_name(guildId, client) -> str: guild = await client.fetch_guild(guildId) return "[Server={0}]".format(guild.name) + async def handle_debug_commands(message, client) -> bool: """Handle debug commands (debug, save, load). Returns True if handled.""" message_content = message.content.lower() - + if message_content == "debug poilau": if message.author.id == 151626081458192384: - all_states = get_all_guild_states() + all_states = database.get_all_guild_states() dump = {} for guild_id, state in all_states.items(): channel_name = await get_guild_name(guild_id, client) @@ -182,43 +91,49 @@ async def handle_debug_commands(message, client) -> bool: "cooldown_until": state["cooldown_until"], "cooldown_remaining": sleeping_time, "self-control": state["self_control"], - "last_updated": state["last_updated"] + "last_updated": state["last_updated"], } await message.author.send( "```json\n{0}```".format(json.dumps(dump, ensure_ascii=False, indent=2)) ) return True - + if message_content == "save poilau": if message.author.id == 151626081458192384: - all_states = get_all_guild_states() - json_str = "```json\n{0}```".format(json.dumps(all_states, ensure_ascii=False, indent=2)) + all_states = database.get_all_guild_states() + json_str = "```json\n{0}```".format( + json.dumps(all_states, ensure_ascii=False, indent=2) + ) await message.author.send("State persisted in SQLite database") await message.author.send(json_str) return True - + if message_content == "load poilau": if message.author.id == 151626081458192384: success, msg = load_rhymes() - all_states = get_all_guild_states() - json_str = "```json\n{0}```".format(json.dumps(all_states, ensure_ascii=False, indent=2)) + all_states = database.get_all_guild_states() + json_str = "```json\n{0}```".format( + json.dumps(all_states, ensure_ascii=False, indent=2) + ) await message.author.send(msg) await message.author.send(json_str) return True - + if message_content == "tg fouras" and message.guild: # Disable cooldown for this server (set to far future) - cooldown_date = datetime.now().replace(hour=0, minute=0, second=0, microsecond=0) + cooldown_date = datetime.now().replace( + hour=0, minute=0, second=0, microsecond=0 + ) cooldown_date = cooldown_date.replace(day=cooldown_date.day + 10000) - update_guild_state( + database.update_guild_state( str(message.guild.id), guild_name=message.guild.name, cooldown_until=cooldown_date.isoformat(), - self_control=2.0 + self_control=2.0, ) await message.channel.send("ok :'(") return True - + return False @@ -226,73 +141,77 @@ async def handle_rhyme_logic(message, client) -> bool: """Main rhyme detection logic. Returns True if rhyme was triggered.""" message_content = message.content.lower() last_word = get_last_word(message_content) - + if message.author != client.user and message.guild and last_word: rhyme = find_rhyme(last_word) guild_id = str(message.guild.id) guild_name = message.guild.name - + if rhyme: - guild_state = get_guild_state(guild_id) - + guild_state = database.get_guild_state(guild_id) + # Update guild name if changed if guild_state["guild_name"] != guild_name: - update_guild_state( + database.update_guild_state( guild_id, guild_name=guild_name, cooldown_until=guild_state["cooldown_until"], - self_control=guild_state["self_control"] + self_control=guild_state["self_control"], ) - + # Check cooldown cooldown_dt = datetime.fromisoformat(guild_state["cooldown_until"]) now_dt = datetime.now() - + if now_dt >= cooldown_dt: self_control = guild_state["self_control"] - + # Probability check if random.random() < self_control: new_self_control = self_control * 0.9 - update_guild_state( + database.update_guild_state( guild_id, guild_name=guild_name, cooldown_until=now_dt.isoformat(), - self_control=new_self_control + self_control=new_self_control, ) return False - + # Calculate new cooldown duration wait_time = random.randint(0, 900) if bool(random.getrandbits(1)): wait_time = random.randint(900, 10800) - + new_cooldown_dt = now_dt.replace(second=0, microsecond=0) - new_cooldown_dt = new_cooldown_dt.replace(minute=new_cooldown_dt.minute + wait_time // 60) - new_cooldown_dt = new_cooldown_dt.replace(hour=new_cooldown_dt.hour + wait_time // 3600) - - update_guild_state( + new_cooldown_dt = new_cooldown_dt.replace( + minute=new_cooldown_dt.minute + wait_time // 60 + ) + new_cooldown_dt = new_cooldown_dt.replace( + hour=new_cooldown_dt.hour + wait_time // 3600 + ) + + database.update_guild_state( guild_id, guild_name=guild_name, cooldown_until=new_cooldown_dt.isoformat(), - self_control=self_control + 1.0 + self_control=self_control + 1.0, ) - + await message.channel.send(rhyme) return True - + return False async def handle_message(message, client) -> bool: """Main entry point for message handling.""" # Initialize database and log file on first run - _ensure_db() + database.ensure_db() _ensure_log_file() - + # Handle debug commands first if await handle_debug_commands(message, client): return True - + # Process rhyme logic - return await handle_rhyme_logic(message, client) \ No newline at end of file + return await handle_rhyme_logic(message, client)