diff --git a/config.json.example b/config.json.example index c7401f8..9c7767e 100644 --- a/config.json.example +++ b/config.json.example @@ -1,20 +1,33 @@ { + "discord": { "token": "", - "stations": ["Hamburg Hbf", "München Hbf", "Köln Hbf", "Amsterdam Centraal"], - "dbf": "https://dbf.finalrewind.org", - "server": , + "server": , "vc": , - "random": true, - "emojis": true, "formatting": "┇", - "announcements": true, - "voice_announcements: [ - { - "enabled": false, - "stations": { - "general": "" - } + "emojis": true + }, + "connections": { + "stations": [ + "" + ], + "blacklist": [], + "min_duration": 5, + "max_duration": null, + "max_wait_time": 6, + "timezone": "Europe/Berlin" + }, + "announcements": { + "enabled": true, + "voice": [ + { + "enabled": false, + "stations": { + "general": "general.aac", } + } ] - "blacklist": [] + }, + "http": { + "user_agent": "Gleiswechsel-Discord-Bot" + } } \ No newline at end of file diff --git a/main.py b/main.py index 2161cd4..53b8c97 100644 --- a/main.py +++ b/main.py @@ -1,5 +1,6 @@ import discord -from src.utils import config, logger +from src.utils import logger +from src.config import config from src.dc.handlers import rename_vc from src.dc.helpers import validate_channel from src.dc.commands import setup_commands @@ -16,22 +17,21 @@ async def on_ready(): if not _bot_initialized: _bot_initialized = True - server_id = config["server"] - server_vc_id = config["vc"] + server_id = config.discord.server + server_vc_id = config.discord.vc channel = validate_channel(bot=bot, server_id=server_id, channel_id=server_vc_id) await rename_vc(bot, voice_channel=channel) else: logger("Reconnected to discord gateway, this wont disturb your current ride") try: - bot.run(config["token"]) + bot.run(config.discord.token) except: logger("Feher peim parsen des tokens", "fatal") ''' TODO - discord status -- config cleanup - multi language support - README ''' \ No newline at end of file diff --git a/src/api/transitous.py b/src/api/transitous.py index ed75d21..0f26fe0 100644 --- a/src/api/transitous.py +++ b/src/api/transitous.py @@ -3,11 +3,12 @@ import random import json from datetime import datetime, timezone -from src.utils import logger, config, get_train_name, convert_iso_string, validate_connection +from src.utils import logger, get_train_name, convert_iso_string, validate_connection +from src.config import config -stations = config["stations"] -blacklist = config["blacklist"] -user_agent = config["http"]["user_agent"] +stations = config.connections.stations +blacklist = config.connections.blacklist +user_agent = config.http.user_agent headers = { "User-Agent": f"{user_agent}" diff --git a/src/config.py b/src/config.py new file mode 100644 index 0000000..160486c --- /dev/null +++ b/src/config.py @@ -0,0 +1,58 @@ +import json +import os +from dataclasses import dataclass, field +from typing import Optional + +@dataclass +class DiscordConfig: + token: str + server: int + vc: int + formatting: str + emojis: bool + +@dataclass +class ConnectionsConfig: + stations: list[str] + blacklist: list[str] + min_duration: int + max_wait_time: int + timezone: str + max_duration: Optional[int] = None + +@dataclass +class VoiceAnnouncementConfig: + enabled: bool + stations: dict[str, str] + +@dataclass +class AnnouncementConfig: + enabled: bool + voice: list[VoiceAnnouncementConfig] + +@dataclass +class HttpConfig: + user_agent: str + +@dataclass +class Config: + discord: DiscordConfig + connections: ConnectionsConfig + announcements: AnnouncementConfig + http: HttpConfig + +def _load_config() -> Config: + with open("config.json", "r") as file: + raw = json.load(file) + + return Config( + discord=DiscordConfig(**raw["discord"]), + connections=ConnectionsConfig(**raw["connections"]), + announcements=AnnouncementConfig( + enabled=raw["announcements"]["enabled"], + voice=[VoiceAnnouncementConfig(**v) for v in raw["announcements"]["voice"]], + ), + http=HttpConfig(**raw["http"]), + ) + +config = _load_config() \ No newline at end of file diff --git a/src/dc/handlers.py b/src/dc/handlers.py index fa318cc..c2c00c9 100644 --- a/src/dc/handlers.py +++ b/src/dc/handlers.py @@ -3,7 +3,8 @@ import asyncio import random from datetime import datetime, timedelta, date -from src.utils import logger, channel_formatting, choose_connection, config, get_sound_path +from src.utils import logger, channel_formatting, choose_connection, get_sound_path +from src.config import config _scheduled_task: asyncio.Task | None = None @@ -62,7 +63,6 @@ async def _schedule_next_transfer(bot: discord.Bot, arrival, voice_channel: disc if wait_seconds > announcement_countdown: wait_until_end_announcement = wait_seconds - announcement_countdown await asyncio.sleep(wait_until_end_announcement) - await asyncio.sleep(5) await announcer("ende", voice_channel, destination) await asyncio.sleep(announcement_countdown) else: @@ -74,8 +74,8 @@ async def _schedule_next_transfer(bot: discord.Bot, arrival, voice_channel: disc async def announcer(announcement: str, voice_channel: discord.VoiceChannel, destination = None): from src.dc.embeds import build_info_embed, build_announcement_embed - announcements_enabled = config.get("announcements", True) - voice_announcement_enabled = config["voice_announcements"][0]["enabled"] + announcements_enabled = config.announcements.enabled + voice_announcement_enabled = config.announcements.voice[0].enabled if announcements_enabled: if len(voice_channel.members) > 0: diff --git a/src/dc/helpers.py b/src/dc/helpers.py index 23ee76f..3efc8e8 100644 --- a/src/dc/helpers.py +++ b/src/dc/helpers.py @@ -4,11 +4,12 @@ from src.utils import logger def validate_channel(bot: discord.bot, server_id: int, channel_id: int): guild = bot.get_guild(server_id) - channel = guild.get_channel(channel_id) if guild is None: logger(f"Es konnte kein Server mit der ID {server_id} gefunden werden", "fatal") + return False + channel = guild.get_channel(channel_id) if not isinstance(channel, discord.VoiceChannel): logger(f"Es konnte kein VC mit der id {channel_id} gefunden werden", "fatal") return False diff --git a/src/utils.py b/src/utils.py index ea0989b..2b48988 100644 --- a/src/utils.py +++ b/src/utils.py @@ -1,9 +1,9 @@ -import json import os import importlib import random from pathlib import Path -from datetime import datetime, timedelta, timezone, date +from datetime import datetime, timedelta, timezone +from src.config import config from zoneinfo import ZoneInfo import src.data.operators as operators @@ -11,8 +11,6 @@ from src.data.emojis import emoji_list _operator_mtime = None -with open("config.json", "r") as file: - config = json.load(file) def logger(msg, log_type="info") -> str: status = log_type.upper() @@ -37,7 +35,7 @@ def validate_connection(start_time: str, end_time: str, station_departure: str) logger(f"Verbindung liegt bereits in der Vergangenheit: {start_dt}", "error") return False - max_wait_time = config.get("max_wait_time", 6) + max_wait_time = config.connections.max_wait_time start_dt = datetime.fromisoformat(start_time.replace("Z", "+00:00")) if start_dt > now + timedelta(hours=max_wait_time): logger(f"Verbindung liegt zu weit in der Zukunft: {start_dt}", "error") @@ -47,13 +45,13 @@ def validate_connection(start_time: str, end_time: str, station_departure: str) trip_duration = (end_dt - station_departure_dt).total_seconds() trip_duration_minutes = str(timedelta(seconds=trip_duration)) - min_duration = config.get("min_duration", 10) + min_duration = config.connections.min_duration min_duration_seconds = min_duration * 60 if trip_duration < min_duration_seconds: logger(f"Verbindung ist mit {trip_duration_minutes} zu kurz (mindestens {min_duration} Minuten gewollt)", "error") return False - max_duration = config.get("max_duration", "") + max_duration = config.connections.max_duration if max_duration: max_duration_seconds = max_duration * 60 if max_duration_seconds < trip_duration: @@ -63,7 +61,7 @@ def validate_connection(start_time: str, end_time: str, station_departure: str) return True def convert_iso_string(isostring) -> str: - timezone = config.get("timezone", "Europe/Berlin") + timezone = config.connections.timezone dt = datetime.fromisoformat(isostring.replace('Z', '+00:00')) dt = dt.astimezone(ZoneInfo(timezone)) @@ -72,9 +70,9 @@ def convert_iso_string(isostring) -> str: return dt.strftime('%H:%M') def channel_formatting(mode: str) -> str: - formatting = config.get("formatting", "") + formatting = config.discord.formatting - if config.get("emojis", True): + if config.discord.formatting: emoji = emoji_list.get(mode) if emoji is None: emoji = emoji_list.get("Fallback") @@ -134,14 +132,13 @@ def get_operator_metadata(agency: str, route_color: str) -> dict: } def get_sound_path(destination) -> str | None: - voice_announcement_config = config["voice_announcements"][0] - voice_stations = voice_announcement_config["stations"] + voice_stations = config.announcements.voice[0].stations if destination in voice_stations: announcement_for = destination else: - general_config = voice_stations.get("general", "") - if not general_config: + general_sound_enabled = voice_stations.get("general", "") + if not general_sound_enabled: return None announcement_for = "general"