transitous rewrite: validate connections

This commit is contained in:
Kaaninchen
2026-08-16 13:11:05 +02:00
parent 1892c734cd
commit 1b0b10dc45
6 changed files with 77 additions and 29 deletions
+1 -2
View File
@@ -31,11 +31,10 @@ except:
''' '''
TODO TODO
- 1024 embed limit - 1024 embed limit
- automatic reload of operators
- discord status - discord status
- text announcements - text announcements
- voice announcements - voice announcements
- improved error handling (retry connection) - improved error handling (retry connection)
- multi language support - multi language support
- random = False - README
''' '''
+37 -17
View File
@@ -1,9 +1,9 @@
import requests import requests
import random import random
import json import json
from datetime import datetime, timezone from datetime import datetime, timezone, timedelta
from src.utils import logger, config, get_train_name, convert_iso_string from src.utils import logger, config, get_train_name, convert_iso_string, validate_connection
stations = config["stations"] stations = config["stations"]
blacklist = config["blacklist"] blacklist = config["blacklist"]
@@ -15,8 +15,9 @@ headers = {
endpoint = "https://api.transitous.org" endpoint = "https://api.transitous.org"
def get_random_stop_id() -> str: def get_random_stop_id() -> str | None:
assigned_station = random.choice(stations) assigned_station = random.choice(stations)
logger(f"Station: {assigned_station}")
req = f"{endpoint}/api/v1/geocode" req = f"{endpoint}/api/v1/geocode"
try: try:
@@ -24,18 +25,30 @@ def get_random_stop_id() -> str:
response.raise_for_status() response.raise_for_status()
data = response.json() data = response.json()
except requests.RequestException as e: except requests.RequestException as e:
logger(f"An error occured while searching for a connection: {e}", "fatal") logger(f"An error occured while searching for a connection: {e}", "error")
return None return None
if response.status_code == 404: if response.status_code == 404:
logger(f"Error finding station '{assigned_station}'") logger(f"Error finding station '{assigned_station}'", "error")
return None
id = []
for entry in data: for entry in data:
if entry.get("type") != "STOP": if entry.get("type") != "STOP":
continue continue
return entry["id"] entry_id = entry.get("id", None)
id.append(entry_id)
def get_random_connection(stop_id: str) -> str: if id is None:
logger(f"Failed to grab ID from '{assigned_station}'", "error")
return None
return random.choice(id)
def get_random_connection(stop_id: str) -> str | None:
if stop_id is None:
return None
now = datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ") now = datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")
print(f"Aktuelle Zeit für Query: {now}") print(f"Aktuelle Zeit für Query: {now}")
cursor = None cursor = None
@@ -58,15 +71,13 @@ def get_random_connection(stop_id: str) -> str:
response.raise_for_status() response.raise_for_status()
data = response.json() data = response.json()
except requests.RequestException as e: except requests.RequestException as e:
print(e) logger(e, "Error")
break break
stop_times = data.get("stopTimes", []) stop_times = data.get("stopTimes", [])
all_stop_times.extend(stop_times) all_stop_times.extend(stop_times)
for entry in stop_times: for entry in stop_times:
dep = entry.get("place", {}).get("departure") or entry.get("place", {}).get("arrival")
print(f"{entry.get('mode')}: {convert_iso_string(dep)}")
trip_id = entry["tripId"] trip_id = entry["tripId"]
if entry["mode"] in blacklist: if entry["mode"] in blacklist:
continue continue
@@ -80,7 +91,7 @@ def get_random_connection(stop_id: str) -> str:
break break
if not trip_ids: if not trip_ids:
logger("Couldn't find any connection", "fatal") logger("Couldn't find any connection", "error")
return None return None
trip_id = random.choice(trip_ids) trip_id = random.choice(trip_ids)
@@ -96,28 +107,36 @@ def get_random_connection(stop_id: str) -> str:
"from_station": from_station "from_station": from_station
} }
def get_trip_details(trip_id: str, from_station: str) -> dict: def get_trip_details(random_connection: dict | None) -> dict | None:
if random_connection is None:
return None
req = f"{endpoint}/api/v2/trip" req = f"{endpoint}/api/v2/trip"
try: try:
response = requests.get(req, params={"tripId": trip_id}, headers=headers) response = requests.get(req, params={"tripId": random_connection["trip_id"]}, headers=headers)
response.raise_for_status() response.raise_for_status()
data = response.json() data = response.json()
except requests.RequestException as e: except requests.RequestException as e:
logger(f"An error occured while trying to get the route details: {e}", "fatal") logger(f"An error occured while trying to get the route details: {e}", "error")
return None return None
legs = data["legs"][0] legs = data["legs"][0]
end_time = legs["endTime"]
from_station = random_connection["from_station"]
display_name = legs["displayName"] display_name = legs["displayName"]
trip_from = legs["tripFrom"]["name"] trip_from = legs["tripFrom"]["name"]
goes_to = legs["tripTo"]["name"] goes_to = legs["tripTo"]["name"]
start_time = legs["startTime"] start_time = legs["startTime"]
end_time = legs["endTime"]
mode = legs["mode"] mode = legs["mode"]
departure = convert_iso_string(start_time) is_valid = validate_connection(start_time, end_time)
if not is_valid:
return None
arrival = convert_iso_string(end_time) arrival = convert_iso_string(end_time)
departure = convert_iso_string(start_time)
train_name = get_train_name(display_name, mode) train_name = get_train_name(display_name, mode)
trip_details = { trip_details = {
@@ -143,5 +162,6 @@ def get_trip_details(trip_id: str, from_station: str) -> dict:
trip_details["departure"] = convert_iso_string(departure_time) trip_details["departure"] = convert_iso_string(departure_time)
trip_details["stops"][goes_to] = arrival trip_details["stops"][goes_to] = arrival
logger(json.dumps(trip_details, indent=4, ensure_ascii=False))
return trip_details return trip_details
+2
View File
@@ -0,0 +1,2 @@
SPDX-FileCopyrightText: 2024 Mathis Brüchert <[email protected]>
SPDX-License-Identifier: CC-BY-SA-4.0
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.3 KiB

+17 -5
View File
@@ -11,24 +11,36 @@ async def rename_vc(bot: discord.Bot, voice_channel, from_scheduler: bool = Fals
if not from_scheduler and _scheduled_task and not _scheduled_task.done(): if not from_scheduler and _scheduled_task and not _scheduled_task.done():
_scheduled_task.cancel() _scheduled_task.cancel()
attempt = 0
max_attempt = 10
trip = choose_connection() trip = choose_connection()
while trip is None and attempt < max_attempt:
attempt += 1
logger(f"Attempt {attempt}: Failed to select route, retrying...", "error")
trip = choose_connection()
if trip is None:
logger(f"Failed to select route after {max_attempt} attempts", "fatal")
return False
arrival = trip["arrival"] arrival = trip["arrival"]
long_name = trip["long_name"] long_name = trip["long_name"]
mode = trip["mode"]
print("-----------------") print("-----------------")
logger(f"Umstieg: {long_name}, Ankunft: {arrival} Uhr") logger(f"Umstieg: {long_name}, Ankunft: {arrival} Uhr")
logger(f"Betreiber: {trip["agency"]}, Typ: {trip["mode"]}") logger(f"Betreiber: {trip["agency"]}, Typ: {mode}")
logger(f"Versuche Namen zu ändern, wenn nichts passiert bin ich im cooldown... (warte bis zu 10min!)") logger(f"Versuche Namen zu ändern, wenn nichts passiert bin ich im cooldown... (warte bis zu 10min!)")
formatting = channel_formatting(trip["mode"]) formatting = channel_formatting(mode)
await voice_channel.edit(name=f"{formatting}{long_name}") await voice_channel.edit(name=f"{formatting}{long_name}")
await voice_channel.set_status(f"Ankunft um {arrival}") await voice_channel.set_status(f"Ankunft um {arrival}")
logger(f"Name geändert!") logger(f"Name geändert!")
_scheduled_task = asyncio.create_task(_schedule_next_transfer(bot, arrival)) _scheduled_task = asyncio.create_task(_schedule_next_transfer(bot, arrival, voice_channel))
async def _schedule_next_transfer(bot, arrival): async def _schedule_next_transfer(bot, arrival, voice_channel):
now = datetime.now() now = datetime.now()
parsed_time = datetime.strptime(arrival, "%H:%M").time() parsed_time = datetime.strptime(arrival, "%H:%M").time()
arrival_dt = datetime.combine(date.today(), parsed_time) arrival_dt = datetime.combine(date.today(), parsed_time)
@@ -44,5 +56,5 @@ async def _schedule_next_transfer(bot, arrival):
await asyncio.sleep(wait_seconds) await asyncio.sleep(wait_seconds)
logger("Zug angekommen, wähle neue Verbindung") logger("Zug angekommen, wähle neue Verbindung")
await rename_vc(bot, from_scheduler=True) await rename_vc(bot, voice_channel, from_scheduler=True)
+20 -5
View File
@@ -1,7 +1,7 @@
import json import json
import os import os
import importlib import importlib
from datetime import datetime, timedelta from datetime import datetime, timedelta, timezone
from zoneinfo import ZoneInfo from zoneinfo import ZoneInfo
import src.data.operators as operators import src.data.operators as operators
@@ -19,14 +19,30 @@ def logger(msg, log_type="info") -> str:
if status == "FATAL": if status == "FATAL":
os._exit(1) os._exit(1)
def choose_connection() -> dict: def choose_connection() -> dict | None:
from src.api import transitous from src.api import transitous
station_id = transitous.get_random_stop_id() station_id = transitous.get_random_stop_id()
connection = transitous.get_random_connection(station_id) connection = transitous.get_random_connection(station_id)
trip = transitous.get_trip_details(connection["trip_id"], connection["from_station"]) trip = transitous.get_trip_details(connection)
return trip return trip
def validate_connection(start_time: str, end_time: str) -> bool:
now = datetime.now(timezone.utc)
start_dt = datetime.fromisoformat(start_time.replace("Z", "+00:00"))
end_dt = datetime.fromisoformat(end_time.replace("Z", "+00:00"))
max_wait_time = config.get("max_wait_time", 6)
if end_dt < now:
logger(f"Verbindung liegt bereits in der Vergangenheit: {start_dt}", "error")
return False
if start_dt > now + timedelta(hours=max_wait_time):
logger(f"Verbindung liegt zu weit in der Zukunft: {start_dt}", "error")
return False
return True
def convert_iso_string(isostring) -> str: def convert_iso_string(isostring) -> str:
timezone = config.get("timezone", "Europe/Berlin") timezone = config.get("timezone", "Europe/Berlin")
dt = datetime.fromisoformat(isostring.replace('Z', '+00:00')) dt = datetime.fromisoformat(isostring.replace('Z', '+00:00'))
@@ -34,7 +50,6 @@ def convert_iso_string(isostring) -> str:
if dt.second >= 30: if dt.second >= 30:
dt += timedelta(minutes=1) dt += timedelta(minutes=1)
return dt.strftime('%H:%M') return dt.strftime('%H:%M')
def channel_formatting(mode: str) -> str: def channel_formatting(mode: str) -> str:
@@ -75,7 +90,7 @@ def _reload_operators_if_changed():
def get_operator_metadata(agency: str, route_color: str) -> dict: def get_operator_metadata(agency: str, route_color: str) -> dict:
_reload_operators_if_changed() _reload_operators_if_changed()
op_data = operators.OPERATOR_ALIASES.get(agency) or operators.OPERATORS.get(agency) or operators.OPERATORS["fallback"] op_data = operators.OPERATOR_ALIASES.get(agency) or operators.OPERATORS.get(agency) or operators.OPERATORS["fallback"]
logo = op_data.get("logo", operators.OPERATORS["fallback"]["logo"]) logo = op_data.get("logo", operators.OPERATORS["fallback"]["logo"])