transitous rewrite: next_stop improvements

This commit is contained in:
Kaaninchen
2026-08-17 12:49:35 +02:00
parent 27f8ef2cbc
commit 59af14e6ff
6 changed files with 39 additions and 747 deletions
+17 -17
View File
@@ -3,7 +3,7 @@ import random
import json
from datetime import datetime, timezone
from src.utils import logger, get_train_name, convert_iso_string, validate_connection
from src.utils import logger, get_train_name, validate_connection, parse_iso
from src.config import config
stations = config.connections.stations
@@ -120,8 +120,8 @@ def get_random_stop_id() -> str | None:
stop_ids_list = list(stop_ids.keys())
stops_string = ", ".join(stop_ids.values())
if len(stop_ids_list) > 1:
logger(f"Station '{assigned_station}' not found, choosing random from available: {stops_string}")
logger(f"No station associated as '{assigned_station}', choosing random from available")
logger(f"No station associated as '{assigned_station}', choosing random from similar named stations:")
logger({stops_string})
chosen_stop_id = random.choice(stop_ids_list)
return chosen_stop_id
@@ -209,8 +209,8 @@ def get_trip_details(random_connection: dict | None) -> dict | None:
start_time = legs["startTime"]
mode = legs["mode"]
arrival = convert_iso_string(end_time)
departure = convert_iso_string(start_time)
arrival_dt = parse_iso(end_time)
departure_dt = parse_iso(start_time)
train_name = get_train_name(display_name, mode)
if train_from == from_station:
@@ -227,27 +227,27 @@ def get_trip_details(random_connection: dict | None) -> dict | None:
"agency": legs["agencyName"],
"route_color": legs.get("routeColor"),
"duration": legs["duration"],
"departure": departure,
"arrival": arrival,
"departure": departure_dt.strftime("%H:%M"),
"arrival": arrival_dt.strftime("%H:%M"),
"mode": mode,
"stops": {}
}
trip_details["stops"][train_from] = departure
departure_time = start_time
trip_details["stops"][train_from] = departure_dt
departure_time_iso = start_time
for stop in legs["intermediateStops"]:
stop_arrival = convert_iso_string(stop["arrival"])
trip_details["stops"][stop["name"]] = stop_arrival # not sure if that actually works but im too tired to question it
stop_arrival_dt = parse_iso(stop["arrival"])
trip_details["stops"][stop["name"]] = stop_arrival_dt
if stop.get("name") == from_station:
departure_time = stop["departure"]
trip_details["departure"] = convert_iso_string(departure_time)
trip_details["stops"][goes_to] = arrival
departure_time_iso = stop["departure"]
trip_details["departure"] = parse_iso(departure_time_iso).strftime("%H:%M")
valid = validate_connection(start_time, end_time, departure_time)
trip_details["stops"][goes_to] = arrival_dt
valid = validate_connection(start_time, end_time, departure_time_iso)
if not valid:
return None
logger(json.dumps(trip_details, indent=4, ensure_ascii=False))
return trip_details
+1 -1
View File
@@ -44,7 +44,7 @@ class Config:
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"]),
+1
View File
@@ -110,6 +110,7 @@ OPERATORS = {
}
OPERATOR_ALIASES = {
"Arverio BAyern GmbH": OPERATORS["Arverio Bayern"],
"DB Regio AG Baden-Württemberg": OPERATORS["db_bawü"],
"DB Regio Stuttgart GmbH": OPERATORS["db_bawü"],
"Arverio Baden-Württemberg": OPERATORS["db_bawü"],
+1 -2
View File
@@ -34,8 +34,7 @@ def build_info_embed() -> discord.Embed:
stops = trip["stops"]
next_stop = get_next_station(trip["stops"])
print(f"next_stop: {next_stop}")
next_stop = get_next_station(trip["stops"], trip["from"])
next_stop_station = None
if next_stop:
next_stop_station = next_stop.get("name")
+19 -29
View File
@@ -10,7 +10,7 @@ import src.data.operators as operators
from src.data.emojis import emoji_list
_operator_mtime = None
LOCAL_TZ = ZoneInfo(config.connections.timezone)
def logger(msg, log_type="info") -> str:
status = log_type.upper()
@@ -60,14 +60,9 @@ def validate_connection(start_time: str, end_time: str, station_departure: str)
return True
def convert_iso_string(isostring) -> str:
timezone = config.connections.timezone
dt = datetime.fromisoformat(isostring.replace('Z', '+00:00'))
dt = dt.astimezone(ZoneInfo(timezone))
if dt.second >= 30:
dt += timedelta(minutes=1)
return dt.strftime('%H:%M')
def parse_iso(iso_str: str) -> datetime:
dt = datetime.fromisoformat(iso_str.replace("Z", "+00:00"))
return dt.astimezone(LOCAL_TZ)
def channel_formatting(mode: str) -> str:
formatting = config.discord.formatting
@@ -80,7 +75,7 @@ def channel_formatting(mode: str) -> str:
return f"{emoji}{formatting}"
def get_train_name(train_name: str, mode: str) -> str:
if mode == "BUS" or mode == "TRAM":
if mode == "BUS" or mode == "TRAM" or train_name.isdigit():
train = f"{mode.capitalize()} {train_name}"
elif "(" in train_name:
train = train_name.split(" (")[0]
@@ -153,24 +148,19 @@ def get_sound_path(destination) -> str | None:
return sound_path
def get_next_station(stops: dict) -> dict | None:
now = datetime.now()
today = datetime.today()
day_offset = 0
previous_time = None
for name, arrival_str in stops.items():
arrival_time = datetime.strptime(arrival_str, "%H:%M").time()
if previous_time is not None and arrival_time < previous_time:
day_offset += 1
arrival_dt = datetime.combine(today + timedelta(days=day_offset), arrival_time)
previous_time = arrival_time
def get_next_station(stops: dict, train_from :str) -> dict | None:
now = datetime.now(LOCAL_TZ)
print(train_from)
for name, arrival_dt in stops.items():
if arrival_dt >= now:
return {"name": name, "arrival": arrival_str}
if name == train_from:
return None
else:
return {
"name": name,
"arrival": arrival_dt.strftime("%H:%M")
}
return None
def format_via_list(stops: dict) -> str:
@@ -187,9 +177,9 @@ def format_stop_list(stops: dict, next_stop: str | None) -> list[tuple[str, str]
for name, stop_arrival in stops.items():
if name == next_stop:
line = f"• __{name} ({stop_arrival} Uhr__)"
line = f"• __{name}__ ({stop_arrival.strftime("%H:%M")} Uhr)"
else:
line = f"{name} ({stop_arrival} Uhr)"
line = f"{name} ({stop_arrival.strftime("%H:%M")} Uhr)"
if field_length + len(line) + 1 > 1024:
route_page_name = "Route" if part == 1 else "Route (Fortsetzung)"