"""Pi-side MQTT bridge. Exposes drive()/stop()/read_sensors() — wire these
up as OpenClaw tool functions once OpenClaw is installed. Sends a heartbeat
so the ESP32 watchdog can tell the Pi is still alive.
"""
import json
import threading
import time

import paho.mqtt.client as mqtt

BROKER = "localhost"
PORT = 1883
HEARTBEAT_INTERVAL_S = 0.2

TOPIC_DRIVE = "robot/cmd/drive"
TOPIC_STOP = "robot/cmd/stop"
TOPIC_HEARTBEAT = "robot/heartbeat"
TOPIC_SENSORS = "robot/sensors/status"

VALID_DIRECTIONS = {"forward", "backward", "left", "right"}
MAX_DURATION_MS = 5000

_client = mqtt.Client()
_latest_sensors = {}
_sensors_lock = threading.Lock()


def _on_message(_client, _userdata, msg):
    if msg.topic == TOPIC_SENSORS:
        with _sensors_lock:
            _latest_sensors.update(json.loads(msg.payload))


def connect():
    _client.on_message = _on_message
    _client.connect(BROKER, PORT)
    _client.subscribe(TOPIC_SENSORS)
    _client.loop_start()
    threading.Thread(target=_heartbeat_loop, daemon=True).start()


def _heartbeat_loop():
    while True:
        _client.publish(TOPIC_HEARTBEAT, "1")
        time.sleep(HEARTBEAT_INTERVAL_S)


def drive(direction, speed=150, duration_ms=1000):
    if direction not in VALID_DIRECTIONS:
        raise ValueError(f"direction must be one of {VALID_DIRECTIONS}")
    speed = max(0, min(255, speed))
    duration_ms = max(0, min(MAX_DURATION_MS, duration_ms))

    _client.publish(TOPIC_DRIVE, json.dumps({"direction": direction, "speed": speed}))
    timer = threading.Timer(duration_ms / 1000, stop)
    timer.start()
    return {"direction": direction, "speed": speed, "duration_ms": duration_ms}


def stop():
    _client.publish(TOPIC_STOP, json.dumps({}))


def read_sensors():
    with _sensors_lock:
        return dict(_latest_sensors)


if __name__ == "__main__":
    import sys

    connect()
    time.sleep(0.3)  # let the MQTT connection establish
    args = sys.argv[1:]

    if not args:
        print("Bridge running. Try: drive('forward'); time.sleep(1); print(read_sensors())")
        while True:
            time.sleep(1)
    elif args[0] == "drive":
        direction = args[1]
        speed = int(args[2]) if len(args) > 2 else 150
        duration_ms = int(args[3]) if len(args) > 3 else 1000
        result = drive(direction, speed, duration_ms)
        time.sleep(duration_ms / 1000)
        print(json.dumps(result))
    elif args[0] == "stop":
        stop()
        print("stopped")
    elif args[0] == "sensors":
        time.sleep(1.0)  # wait for at least one periodic sensor publish
        print(json.dumps(read_sensors()))
    else:
        print(f"unknown command: {args[0]}", file=sys.stderr)
        sys.exit(1)
