#!/usr/bin/env python3
"""
Fleet telemetry load generator.

Publishes OMNI-shaped frames straight to the MQTT broker, which is the path a real device takes:
EMQX -> iot-bridge-service -> Kafka -> fleet-service / booking-service. Nothing is faked further
down the chain, so what this measures is what production would do.

    python3 mqtt_load.py --devices 2000 --hz 0.333 --seconds 60

  --devices  how many distinct vehicles report
  --hz       frames per second per device (0.333 = every 3 s)
  --seconds  how long to run

It speaks MQTT 3.1.1 over a plain socket. That is deliberate: no pip install, so it runs on a
bare box or in CI without a virtualenv.
"""

import argparse
import json
import random
import socket
import struct
import sys
import time


def encode_remaining_length(n: int) -> bytes:
    out = bytearray()
    while True:
        digit = n % 128
        n //= 128
        if n > 0:
            digit |= 0x80
        out.append(digit)
        if n == 0:
            return bytes(out)


def encode_string(s: str) -> bytes:
    raw = s.encode()
    return struct.pack(">H", len(raw)) + raw


class MqttClient:
    """The two packets a publisher needs: CONNECT and PUBLISH at QoS 0."""

    def __init__(self, host: str, port: int, client_id: str):
        self.sock = socket.create_connection((host, port), timeout=10)
        self.sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
        payload = encode_string(client_id)
        variable = encode_string("MQTT") + bytes([4, 0x02]) + struct.pack(">H", 60)
        body = variable + payload
        self.sock.sendall(bytes([0x10]) + encode_remaining_length(len(body)) + body)
        # CONNACK is 4 bytes; a broker that refuses us says so in the last one.
        ack = self.sock.recv(4)
        if len(ack) < 4 or ack[3] != 0:
            raise RuntimeError(f"broker refused the connection: {ack!r}")

    def publish(self, topic: str, payload: bytes) -> None:
        body = encode_string(topic) + payload
        self.sock.sendall(bytes([0x30]) + encode_remaining_length(len(body)) + body)

    def close(self) -> None:
        try:
            self.sock.sendall(bytes([0xE0, 0x00]))   # DISCONNECT
        finally:
            self.sock.close()


def main() -> int:
    ap = argparse.ArgumentParser()
    ap.add_argument("--host", default="localhost")
    ap.add_argument("--port", type=int, default=1883)
    ap.add_argument("--devices", type=int, default=2000)
    ap.add_argument("--hz", type=float, default=1 / 3)
    ap.add_argument("--seconds", type=int, default=60)
    ap.add_argument("--prefix", default="LOAD")
    ap.add_argument("--profile", choices=["omni", "generic"], default="omni",
                    help="which device protocol to speak; 'generic' is the vendor-neutral one")
    ap.add_argument("--connections", type=int, default=8,
                    help="publisher sockets; real devices hold one each, this fans out the load")
    args = ap.parse_args()

    clients = [MqttClient(args.host, args.port, f"loadgen-{i}-{random.randint(0, 1 << 30)}")
               for i in range(args.connections)]
    print(f"connected {len(clients)} publishers to {args.host}:{args.port}")

    interval = 1.0 / args.hz
    target_rate = args.devices * args.hz
    print(f"{args.devices} devices x {args.hz:.3f} Hz = {target_rate:.0f} frames/sec "
          f"for {args.seconds}s (~{int(target_rate * args.seconds):,} frames)")

    # Each device keeps its own position so the frames look like movement rather than noise.
    lat = [24.60 + (d % 200) * 0.001 for d in range(args.devices)]
    lng = [46.60 + (d // 200) * 0.001 for d in range(args.devices)]

    started = time.time()
    deadline = started + args.seconds
    sent = 0
    errors = 0
    next_tick = started
    tick = 0

    while time.time() < deadline:
        tick += 1
        tick_start = time.time()
        for d in range(args.devices):
            lat[d] += random.uniform(-0.00008, 0.00012)
            lng[d] += random.uniform(-0.00008, 0.00012)
            # A digits-only prefix is treated as an IMEI base, so the generator can drive a real
            # registered fleet (IMEIs are 15-17 digits) rather than only unknown devices.
            device = (str(int(args.prefix) + d) if args.prefix.isdigit()
                      else f"{args.prefix}-{d:05d}")
            try:
                client = clients[d % len(clients)]
                if args.profile == "generic":
                    # The vendor-neutral profile: one topic, one flat body. Anything that can be
                    # pointed at an MQTT topic with a templated payload speaks this.
                    body = {"lat": round(lat[d], 6), "lng": round(lng[d], 6)}
                    if tick % 10 == 0:
                        body["battery"] = 40 + (d % 60)
                        body["signal"] = 20 + (d % 10)
                    client.publish(f"beeb/{device}/up", json.dumps(body).encode())
                    sent += 1
                else:
                    # Real OMNI frames, not a convenient JSON of our own: the point of driving the
                    # broker is that every layer below behaves as it will in production, and the
                    # bridge drops anything its codec does not recognise.
                    #   location  -> om/client/data/location/{IMEI}
                    #   heartbeat -> om/client/req/heartbeat/{IMEI}
                    # Coordinates are unsigned magnitudes plus a hemisphere letter (2.4.1).
                    location = json.dumps({
                        "location": [{
                            "LAT": f"{abs(lat[d]):.6f}", "GEO_NS": "N",
                            "LNG": f"{abs(lng[d]):.6f}", "GEO_EW": "E",
                        }]
                    }).encode()
                    client.publish(f"om/client/data/location/{device}", location)
                    sent += 1
                    # A heartbeat every tenth round: the real device sends them far less often
                    # than positions, and sending both every time overstates the battery load.
                    if tick % 10 == 0:
                        heartbeat = json.dumps({
                            "heartbeat": {"sysSoc": 40 + (d % 60), "CSQ": 20 + (d % 10),
                                          "lockSta": 1},
                            "operationId": tick,
                        }).encode()
                        client.publish(f"om/client/req/heartbeat/{device}", heartbeat)
                        sent += 1
            except OSError:
                errors += 1
        # Pace the rounds so the offered rate is the requested one rather than "as fast as the
        # socket will take", which would measure the loopback instead of the platform.
        next_tick += interval
        sleep = next_tick - time.time()
        if sleep > 0:
            time.sleep(sleep)
        elif time.time() - tick_start > interval * 2:
            print(f"  behind schedule: a round took {time.time() - tick_start:.2f}s "
                  f"(budget {interval:.2f}s)")

    elapsed = time.time() - started
    for c in clients:
        c.close()
    print(f"sent {sent:,} frames in {elapsed:.1f}s = {sent / elapsed:,.0f} frames/sec"
          + (f", {errors} socket errors" if errors else ""))
    return 0


if __name__ == "__main__":
    sys.exit(main())
