"""Standalone desktop navigation lab. Synthetic robot, NOT hardware evidence.
Run: python3 lab-kit.py --output lab-results
Faults: --fault hazard | stale | stalled | blocked
Writes all sample rows and an SVG trace; no third-party packages.
"""
import argparse
import csv
import json
from collections import deque
from math import atan2, cos, hypot, pi, sin
from pathlib import Path

CELL_M = 0.20
WIDTH, HEIGHT = 7, 7
DT = 0.05
TRACK_M = 0.10
ROBOT_RADIUS_M, MARGIN_M = 0.04, 0.02
START, GOAL = (1, 1), (5, 5)
OBSTACLES = {(3, 2), (3, 3), (3, 4)}


def inflate(blocked, radius_m):
    """Conservatively test centre distance to each occupied square, and walls."""
    result = set(blocked)
    for x in range(WIDTH):
        for y in range(HEIGHT):
            px, py = (x + 0.5) * CELL_M, (y + 0.5) * CELL_M
            if min(px, py, WIDTH * CELL_M - px, HEIGHT * CELL_M - py) <= radius_m:
                result.add((x, y))
            for bx, by in blocked:
                dx = max(bx * CELL_M - px, 0, px - (bx + 1) * CELL_M)
                dy = max(by * CELL_M - py, 0, py - (by + 1) * CELL_M)
                if hypot(dx, dy) <= radius_m:
                    result.add((x, y))
    return result


def bfs(blocked, start, goal):
    def free(p):
        return 0 <= p[0] < WIDTH and 0 <= p[1] < HEIGHT and p not in blocked
    if not free(start) or not free(goal):
        return None
    queue, parents = deque([start]), {start: None}
    while queue:
        current = queue.popleft()
        if current == goal:
            path = []
            while current is not None:
                path.append(current)
                current = parents[current]
            return path[::-1]
        x, y = current
        for nxt in [(x + 1, y), (x - 1, y), (x, y + 1), (x, y - 1)]:
            if free(nxt) and nxt not in parents:
                parents[nxt] = current
                queue.append(nxt)
    return None


def advance(pose, dl, dr):
    x, y, theta = pose
    ds, turn = (dl + dr) / 2, (dr - dl) / TRACK_M
    return (x + ds * cos(theta + turn / 2),
            y + ds * sin(theta + turn / 2), atan2(sin(theta + turn), cos(theta + turn)))


def collision(pose, blocked):
    x, y, _ = pose
    radius = ROBOT_RADIUS_M + MARGIN_M
    if min(x, y, WIDTH * CELL_M - x, HEIGHT * CELL_M - y) <= radius:
        return True
    for bx, by in blocked:
        dx = max(bx * CELL_M - x, 0, x - (bx + 1) * CELL_M)
        dy = max(by * CELL_M - y, 0, y - (by + 1) * CELL_M)
        if hypot(dx, dy) <= radius:
            return True
    return False


def simulate(fault="none", left_scale=1.0):
    if not 0.5 <= left_scale <= 1.5:
        raise ValueError("left_scale must be in [0.5,1.5]")
    blocked = set(OBSTACLES)
    if fault == "blocked":
        blocked.update((4, y) for y in range(HEIGHT))
    route = bfs(inflate(blocked, ROBOT_RADIUS_M + MARGIN_M), START, GOAL)
    pose = ((START[0] + 0.5) * CELL_M, (START[1] + 0.5) * CELL_M, 0.0)
    estimate = pose
    rows = []
    if route is None:
        return {"reason": "no_route", "arrived": False, "evidence": "synthetic simulation",
                "elapsed_s": 0.0, "error_m": None, "route": [], "left_scale": left_scale}, rows, blocked
    targets = [((x + 0.5) * CELL_M, (y + 0.5) * CELL_M) for x, y in route[1:]]
    index, reason, waypoint_started = 0, "timeout", 0.0
    last_movement = 0.0
    for sample in range(1201):
        elapsed = sample * DT
        if fault == "hazard" and elapsed >= 2:
            reason = "hazard"; break
        if fault == "stale" and elapsed >= 2:
            reason = "invalid_data"; break
        if elapsed >= 60:
            reason = "timeout"; break
        if elapsed - waypoint_started >= 15:
            reason = "waypoint_timeout"; break
        if elapsed - last_movement > 1.5:
            reason = "encoder_stall"; break
        tx, ty = targets[index]
        dx, dy = tx - estimate[0], ty - estimate[1]
        distance = hypot(dx, dy)
        if distance <= 0.02:
            index += 1
            if index == len(targets):
                reason = "estimated_arrival"; break
            waypoint_started = elapsed
            continue
        heading_error = atan2(sin(atan2(dy, dx) - estimate[2]), cos(atan2(dy, dx) - estimate[2]))
        omega = max(-0.7, min(0.7, 2 * heading_error))
        speed = 0.0 if abs(heading_error) > 8 * pi / 180 else min(0.10, distance)
        vl, vr = speed - omega * TRACK_M / 2, speed + omega * TRACK_M / 2
        dl, dr = vl * DT, vr * DT
        if fault == "stalled" and elapsed >= 2:
            dl = dr = 0.0
        candidate = advance(pose, dl, dr)
        if collision(candidate, blocked):
            reason = "clearance_guard"; break
        pose = candidate
        estimate = advance(estimate, dl * left_scale, dr)
        if abs(dl) + abs(dr) > 0:
            last_movement = elapsed
        rows.append({"elapsed_s": elapsed, "waypoint": index, "x_m": pose[0], "y_m": pose[1],
                     "estimate_x_m": estimate[0], "estimate_y_m": estimate[1],
                     "heading_rad": pose[2], "left_speed_m_s": vl, "right_speed_m_s": vr})
    goal = ((GOAL[0] + 0.5) * CELL_M, (GOAL[1] + 0.5) * CELL_M)
    error = hypot(pose[0] - goal[0], pose[1] - goal[1])
    result = {"reason": reason, "arrived": reason == "estimated_arrival" and error <= 0.10,
              "evidence": "synthetic simulation, not measured hardware", "elapsed_s": elapsed,
              "error_m": error, "route": route, "left_scale": left_scale,
              "independent_truth": "simulator true pose, not odometry", "fault": fault}
    return result, rows, blocked


def write_trace(path, rows, blocked, route):
    scale = 400 / (WIDTH * CELL_M)
    convert = lambda x, y: (30 + x * scale, 450 - y * scale)
    shapes = ['<svg xmlns="http://www.w3.org/2000/svg" width="480" height="500" viewBox="0 0 480 500">',
              '<rect width="480" height="500" fill="#faf8f0"/>',
              '<text x="30" y="25" font-family="sans-serif" font-size="16">Synthetic navigation trace — metres</text>']
    for x, y in blocked:
        px, py = convert(x * CELL_M, (y + 1) * CELL_M)
        shapes.append(f'<rect x="{px}" y="{py}" width="{CELL_M * scale}" height="{CELL_M * scale}" fill="#777"/>')
    if route:
        points = ' '.join(f'{px},{py}' for px, py in (convert((x + .5) * CELL_M, (y + .5) * CELL_M) for x, y in route))
        shapes.append(f'<polyline points="{points}" fill="none" stroke="#b37b21" stroke-width="2" stroke-dasharray="5 5"/>')
    for xkey, ykey, colour in [('x_m', 'y_m', '#18715a'), ('estimate_x_m', 'estimate_y_m', '#2952aa')]:
        points = ' '.join(f'{px},{py}' for px, py in (convert(r[xkey], r[ykey]) for r in rows))
        shapes.append(f'<polyline points="{points}" fill="none" stroke="{colour}" stroke-width="2"/>')
    shapes.extend(['<text x="30" y="480" font-family="sans-serif" font-size="12">Green: true pose · Blue: odometry · Gold: planned cells</text>', '</svg>'])
    path.write_text('\n'.join(shapes))


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output', type=Path, default=Path('lab-results'))
    parser.add_argument('--fault', choices=['none', 'hazard', 'stale', 'stalled', 'blocked'], default='none')
    parser.add_argument('--left-scale', type=float, default=1.0)
    args = parser.parse_args()
    result, rows, blocked = simulate(args.fault, args.left_scale)
    args.output.mkdir(parents=True, exist_ok=True)
    (args.output / 'summary.json').write_text(json.dumps(result, indent=2) + '\n')
    with (args.output / 'trace.csv').open('w', newline='') as handle:
        fields = ['elapsed_s', 'waypoint', 'x_m', 'y_m', 'estimate_x_m', 'estimate_y_m',
                  'heading_rad', 'left_speed_m_s', 'right_speed_m_s']
        writer = csv.DictWriter(handle, fieldnames=fields)
        writer.writeheader(); writer.writerows(rows)
    write_trace(args.output / 'trace.svg', rows, blocked, result['route'])
    print(json.dumps(result, indent=2))

if __name__ == '__main__':
    main()
