from __future__ import annotations import copy import json import math from pathlib import Path from typing import Any from . import ( AccelerationProfile, Color, MTSegment, RepeatedSMSegment, SMSegment, STSegment, make_mileage_time_scheme, make_repeated_speed_mileage_time_scheme, make_speed_mileage_scheme, make_speed_time_scheme, ) SCHEMA_VERSION = 1 SCHEME_KINDS = ("ST", "SM", "MT", "RSMT") ACCELERATION_PROFILES = ("instant", "smooth") def _sm(speed: float, mileage_m: int) -> SMSegment: segment = SMSegment() segment.speed_m_s = speed segment.mileage_from_start_m = mileage_m return segment def _mt(mileage_m: int, time_s: int) -> MTSegment: segment = MTSegment() segment.mileage_to_travel_this_segment_m = mileage_m segment.time_since_start_s = time_s return segment def _st(speed: float, time_s: int) -> STSegment: segment = STSegment() segment.speed_m_s = speed segment.time_since_start_s = time_s return segment def _rsmt(time_s: int, segments: list[SMSegment]) -> RepeatedSMSegment: segment = RepeatedSMSegment() segment.time_since_start_s = time_s segment.speed_mileage_segments = segments return segment def default_draft(kind: str = "ST") -> dict[str, Any]: if kind == "SM": segments: list[dict[str, Any]] = [ {"mileage_m": 0, "speed_m_s": 1.0}, {"mileage_m": 10, "speed_m_s": 2.8}, {"mileage_m": 24, "speed_m_s": 1.2}, {"mileage_m": 40, "speed_m_s": 2.0}, ] elif kind == "MT": segments = [ {"time_s": 0, "mileage_m": 12}, {"time_s": 6, "mileage_m": 20}, {"time_s": 18, "mileage_m": 10}, {"time_s": 28, "mileage_m": 1}, ] elif kind == "RSMT": segments = [ { "time_s": 0, "sub_segments": [ {"mileage_m": 0, "speed_m_s": 1.0}, {"mileage_m": 10, "speed_m_s": 3.0}, {"mileage_m": 24, "speed_m_s": 1.2}, ], }, { "time_s": 18, "sub_segments": [ {"mileage_m": 0, "speed_m_s": 2.2}, {"mileage_m": 8, "speed_m_s": 4.0}, {"mileage_m": 20, "speed_m_s": 1.0}, ], }, { "time_s": 42, "sub_segments": [ {"mileage_m": 0, "speed_m_s": 1.4}, {"mileage_m": 12, "speed_m_s": 2.4}, {"mileage_m": 24, "speed_m_s": 1.4}, ], }, ] else: kind = "ST" segments = [ {"time_s": 0, "speed_m_s": 1.0}, {"time_s": 8, "speed_m_s": 3.2}, {"time_s": 18, "speed_m_s": 1.4}, {"time_s": 30, "speed_m_s": 2.4}, {"time_s": 42, "speed_m_s": 0.8}, ] return { "version": SCHEMA_VERSION, "id": 1, "kind": kind, "acceleration_profile": "smooth", "color": [0, 255, 0], "segments": segments, } def validation_error(draft: dict[str, Any]) -> str | None: try: normalize_draft(draft) except ValueError as exc: return str(exc) return None def normalize_draft(draft: dict[str, Any]) -> dict[str, Any]: if not isinstance(draft, dict): raise ValueError("scheme JSON must be an object") version = _integer(draft.get("version", SCHEMA_VERSION), "version") if version != SCHEMA_VERSION: raise ValueError(f"unsupported scheme version {version}") kind = str(draft.get("kind", "ST")) if kind not in SCHEME_KINDS: raise ValueError(f"unsupported scheme kind {kind!r}") profile = str(draft.get("acceleration_profile", "smooth")) if profile not in ACCELERATION_PROFILES: raise ValueError(f"unsupported acceleration profile {profile!r}") normalized = { "version": SCHEMA_VERSION, "id": _bounded_int(draft.get("id", 1), "id", 0, 255), "kind": kind, "acceleration_profile": profile, "color": _color_list(draft.get("color", [0, 255, 0])), "segments": _normalize_segments(kind, draft.get("segments", [])), } return normalized def clone_draft(draft: dict[str, Any]) -> dict[str, Any]: return copy.deepcopy(draft) def load_draft(path: str | Path) -> dict[str, Any]: with Path(path).open("r", encoding="utf-8") as stream: return normalize_draft(json.load(stream)) def save_draft(path: str | Path, draft: dict[str, Any]) -> dict[str, Any]: normalized = normalize_draft(draft) with Path(path).open("w", encoding="utf-8") as stream: json.dump(normalized, stream, indent=2) stream.write("\n") return normalized def draft_to_scheme(draft: dict[str, Any]): normalized = normalize_draft(draft) scheme_id = int(normalized["id"]) color = Color(*normalized["color"]) profile = ( AccelerationProfile.smooth if normalized["acceleration_profile"] == "smooth" else AccelerationProfile.instant ) kind = normalized["kind"] segments = normalized["segments"] if kind == "SM": return make_speed_mileage_scheme( scheme_id, color, profile, [_sm(float(segment["speed_m_s"]), int(segment["mileage_m"])) for segment in segments], ) if kind == "MT": return make_mileage_time_scheme( scheme_id, color, profile, [_mt(int(segment["mileage_m"]), int(segment["time_s"])) for segment in segments], ) if kind == "RSMT": return make_repeated_speed_mileage_time_scheme( scheme_id, color, profile, [ _rsmt( int(segment["time_s"]), [ _sm(float(sub_segment["speed_m_s"]), int(sub_segment["mileage_m"])) for sub_segment in segment["sub_segments"] ], ) for segment in segments ], ) return make_speed_time_scheme( scheme_id, color, profile, [_st(float(segment["speed_m_s"]), int(segment["time_s"])) for segment in segments], ) def _normalize_segments(kind: str, value: Any) -> list[dict[str, Any]]: if not isinstance(value, list): raise ValueError("segments must be a list") if kind == "ST": segments = [ { "time_s": _bounded_int(segment.get("time_s"), "time_s", 0, 65535), "speed_m_s": _nonnegative_float(segment.get("speed_m_s"), "speed_m_s"), } for segment in _objects(value, "segments") ] _require_nonempty(segments, "ST") segments.sort(key=lambda segment: segment["time_s"]) _require_first_zero(segments[0]["time_s"], "first ST time_s") _require_strictly_increasing(segments, "time_s") return segments if kind == "SM": segments = [ { "mileage_m": _bounded_int(segment.get("mileage_m"), "mileage_m", 0, 65535), "speed_m_s": _nonnegative_float(segment.get("speed_m_s"), "speed_m_s"), } for segment in _objects(value, "segments") ] _require_nonempty(segments, "SM") segments.sort(key=lambda segment: segment["mileage_m"]) _require_first_zero(segments[0]["mileage_m"], "first SM mileage_m") _require_strictly_increasing(segments, "mileage_m") return segments if kind == "MT": segments = [ { "time_s": _bounded_int(segment.get("time_s"), "time_s", 0, 65535), "mileage_m": _bounded_int(segment.get("mileage_m"), "mileage_m", 0, 65535), } for segment in _objects(value, "segments") ] if len(segments) < 2: raise ValueError("MT needs at least two segments") segments.sort(key=lambda segment: segment["time_s"]) _require_first_zero(segments[0]["time_s"], "first MT time_s") _require_strictly_increasing(segments, "time_s") for index, segment in enumerate(segments[:-1]): if segment["mileage_m"] <= 0: raise ValueError(f"MT segment {index} mileage_m must be greater than 0") return segments segments = [] for index, segment in enumerate(_objects(value, "segments")): sub_segments = [ { "mileage_m": _bounded_int(sub_segment.get("mileage_m"), "mileage_m", 0, 65535), "speed_m_s": _nonnegative_float(sub_segment.get("speed_m_s"), "speed_m_s"), } for sub_segment in _objects(segment.get("sub_segments", []), f"RSMT segment {index} sub_segments") ] _require_nonempty(sub_segments, f"RSMT segment {index}") sub_segments.sort(key=lambda sub_segment: sub_segment["mileage_m"]) _require_first_zero(sub_segments[0]["mileage_m"], f"first RSMT segment {index} mileage_m") _require_strictly_increasing(sub_segments, "mileage_m") segments.append( { "time_s": _bounded_int(segment.get("time_s"), "time_s", 0, 65535), "sub_segments": sub_segments, } ) _require_nonempty(segments, "RSMT") segments.sort(key=lambda segment: segment["time_s"]) _require_first_zero(segments[0]["time_s"], "first RSMT time_s") _require_strictly_increasing(segments, "time_s") return segments def _objects(value: Any, name: str) -> list[dict[str, Any]]: if not isinstance(value, list): raise ValueError(f"{name} must be a list") if not all(isinstance(item, dict) for item in value): raise ValueError(f"{name} entries must be objects") return value def _require_nonempty(segments: list[dict[str, Any]], kind: str) -> None: if not segments: raise ValueError(f"{kind} needs at least one segment") def _require_first_zero(value: int, name: str) -> None: if value != 0: raise ValueError(f"{name} must be 0") def _require_strictly_increasing(segments: list[dict[str, Any]], field: str) -> None: for index in range(1, len(segments)): if int(segments[index][field]) <= int(segments[index - 1][field]): raise ValueError(f"{field} must be strictly increasing") def _integer(value: Any, name: str) -> int: if isinstance(value, bool): raise ValueError(f"{name} must be an integer") try: result = int(value) except (TypeError, ValueError) as exc: raise ValueError(f"{name} must be an integer") from exc if result != value and not (isinstance(value, float) and value.is_integer()): raise ValueError(f"{name} must be an integer") return result def _bounded_int(value: Any, name: str, minimum: int, maximum: int) -> int: result = _integer(value, name) if result < minimum or result > maximum: raise ValueError(f"{name} must be between {minimum} and {maximum}") return result def _nonnegative_float(value: Any, name: str) -> float: try: result = float(value) except (TypeError, ValueError) as exc: raise ValueError(f"{name} must be a number") from exc if not math.isfinite(result) or result < 0.0: raise ValueError(f"{name} must be a finite non-negative number") return result def _color_list(value: Any) -> list[int]: if not isinstance(value, (list, tuple)) or len(value) < 3: raise ValueError("color must contain red, green, and blue values") return [ _bounded_int(value[0], "color red", 0, 255), _bounded_int(value[1], "color green", 0, 255), _bounded_int(value[2], "color blue", 0, 255), ]