mirror of
https://github.com/kaaninchen/Gleiswechsel.git
synced 2026-09-17 16:52:47 +00:00
transitous rewrite: refactor config
This commit is contained in:
@@ -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}"
|
||||
|
||||
@@ -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()
|
||||
+4
-4
@@ -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:
|
||||
|
||||
+2
-1
@@ -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
|
||||
|
||||
+11
-14
@@ -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"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user