diff --git a/app/config.py b/app/config.py index e41f30c..8c58e5d 100644 --- a/app/config.py +++ b/app/config.py @@ -13,12 +13,18 @@ class StationConfig: exclude_destinations: List[str] = field(default_factory=list) +@dataclass +class ProfileConfig: + name: str + stations: List[StationConfig] = field(default_factory=list) + + @dataclass class AppConfig: refresh_seconds: int cache_seconds: int departures_limit: int - stations: List[StationConfig] + profiles: List[ProfileConfig] def load_config(path: str = None) -> AppConfig: @@ -26,18 +32,21 @@ def load_config(path: str = None) -> AppConfig: with open(path, "r", encoding="utf-8") as f: raw = yaml.safe_load(f) - stations = [ - StationConfig( - name=s["name"], - type=s["type"].upper(), - exclude_destinations=s.get("exclude_destinations", []) or [], - ) - for s in raw.get("stations", []) - ] + profiles = [] + for p in raw.get("profiles", []): + stations = [ + StationConfig( + name=s["name"], + type=s["type"].upper(), + exclude_destinations=s.get("exclude_destinations", []) or [], + ) + for s in p.get("stations", []) + ] + profiles.append(ProfileConfig(name=p["name"], stations=stations)) return AppConfig( refresh_seconds=int(raw.get("refresh_seconds", 60)), cache_seconds=int(raw.get("cache_seconds", 20)), departures_limit=int(raw.get("departures_limit", 10)), - stations=stations, + profiles=profiles, ) diff --git a/app/main.py b/app/main.py index 69ef276..2936e5e 100644 --- a/app/main.py +++ b/app/main.py @@ -12,7 +12,7 @@ from fastapi.staticfiles import StaticFiles from mvg import MvgApi, TransportType -from app.config import load_config, StationConfig +from app.config import load_config, StationConfig, ProfileConfig logging.basicConfig(level=logging.INFO) logger = logging.getLogger("mvg-departures") @@ -98,27 +98,40 @@ def fetch_departures_for_station(station_cfg: StationConfig) -> List[Dict[str, A return result -def get_all_departures() -> List[Dict[str, Any]]: +def get_profile_by_name(profile_name: str) -> ProfileConfig: + for profile in config.profiles: + if profile.name == profile_name: + return profile + return config.profiles[0] if config.profiles else None + + +def get_departures_for_profile(profile: ProfileConfig) -> List[Dict[str, Any]]: all_deps: List[Dict[str, Any]] = [] - for station_cfg in config.stations: + for station_cfg in profile.stations: all_deps.extend(fetch_departures_for_station(station_cfg)) all_deps.sort(key=lambda d: d["time_epoch"] or 0) return all_deps @app.get("/api/departures") -def api_departures(): - return JSONResponse(content={"departures": get_all_departures()}) +def api_departures(profile: str = None): + active_profile = get_profile_by_name(profile) if profile else config.profiles[0] if config.profiles else None + if not active_profile: + return JSONResponse(content={"departures": []}) + return JSONResponse(content={"departures": get_departures_for_profile(active_profile)}) @app.get("/") -def index(request: Request): - departures = get_all_departures() +def index(request: Request, profile: str = None): + active_profile = get_profile_by_name(profile) if profile else config.profiles[0] if config.profiles else None + departures = get_departures_for_profile(active_profile) if active_profile else [] return templates.TemplateResponse( request, "index.html", { "departures": departures, + "profiles": config.profiles, + "active_profile": active_profile, "refresh_seconds": config.refresh_seconds, "generated_at": datetime.now().strftime("%H:%M:%S"), }, diff --git a/app/templates/index.html b/app/templates/index.html index 8de06dd..573348c 100644 --- a/app/templates/index.html +++ b/app/templates/index.html @@ -5,11 +5,23 @@