#!/usr/bin/env python3
"""Subscribe to the GFEX futures L2 streams.

The server sends compact binary records, not JSON.  This example decodes the
three GFEX subjects: L2 snapshots, volume statistics, and top-ten orders.

Example::

    python3 examples/gfex_l2_subscriber.py \
        --url nats://nats.example.com:4222 \
        --user "$NATS_USER" --password "$NATS_PASSWORD" \
        --all
"""

from __future__ import annotations

import argparse
import asyncio
import json
import os
import struct
import sys
from typing import Any

from nats.aio.client import Client as NATS


# ``<`` means little-endian and standard sizes; no native alignment is used.
GFEX_L2 = struct.Struct(
    "<32s3i4d2qd3q5d2q2d5d5q5i5d5q5i4d"
)
GFEX_L2_SIZE = GFEX_L2.size
if GFEX_L2_SIZE != 428:
    raise RuntimeError(f"unexpected GFEX L2 ABI size: {GFEX_L2_SIZE}")

GFEX_QTY = struct.Struct("<32s3i5d20q")
GFEX_QTY_SIZE = GFEX_QTY.size
if GFEX_QTY_SIZE != 244:
    raise RuntimeError(f"unexpected GFEX quantity ABI size: {GFEX_QTY_SIZE}")

GFEX_ORDER = struct.Struct("<32s3id10qd10q")
GFEX_ORDER_SIZE = GFEX_ORDER.size
if GFEX_ORDER_SIZE != 220:
    raise RuntimeError(f"unexpected GFEX order ABI size: {GFEX_ORDER_SIZE}")


def _text(value: bytes) -> str:
    return value.split(b"\0", 1)[0].decode("ascii", errors="replace")


def _time_text(value: int) -> str:
    """Convert HHMMSSmmm (the MdlService integer representation)."""
    milliseconds = value % 1000
    seconds = (value // 1000) % 100
    minutes = (value // 100000) % 100
    hours = value // 10000000
    return f"{hours:02d}:{minutes:02d}:{seconds:02d}.{milliseconds:03d}"


def decode_gfex_l2(payload: bytes) -> dict[str, Any]:
    """Decode one ``gfex.fut_l2.*`` payload to a Python dict."""
    if len(payload) != GFEX_L2_SIZE:
        raise ValueError(
            f"expected {GFEX_L2_SIZE} bytes for GfexFutureL2, got {len(payload)}"
        )

    values = list(GFEX_L2.unpack(payload))
    symbol = _text(values.pop(0))
    action_day, trading_day, update_time = values[:3]
    del values[:3]

    scalar_names = (
        "last_price",
        "high_price",
        "low_price",
        "open_price",
        "last_volume",
        "volume",
        "turnover",
        "open_interest",
        "pre_open_interest",
        "open_interest_chg",
        "average_price",
        "close_price",
        "settlement_price",
        "pre_settlement_price",
        "pre_close_price",
        "buy_volume",
        "sell_volume",
        "avg_buy_price",
        "avg_sell_price",
    )
    scalars = dict(zip(scalar_names, values[: len(scalar_names)]))
    del values[: len(scalar_names)]

    bid_price = values[:5]
    del values[:5]
    bid_vol = values[:5]
    del values[:5]
    bid_der_vol = values[:5]
    del values[:5]
    ask_price = values[:5]
    del values[:5]
    ask_vol = values[:5]
    del values[:5]
    ask_der_vol = values[:5]
    del values[:5]
    limit_up, limit_down, life_high, life_low = values

    return {
        "symbol": symbol,
        "action_day": action_day,
        "trading_day": trading_day,
        "update_time": _time_text(update_time),
        **scalars,
        "bid": [
            {
                "price": bid_price[i],
                "volume": bid_vol[i],
                "derived_volume": bid_der_vol[i],
            }
            for i in range(5)
        ],
        "ask": [
            {
                "price": ask_price[i],
                "volume": ask_vol[i],
                "derived_volume": ask_der_vol[i],
            }
            for i in range(5)
        ],
        "limit_up": limit_up,
        "limit_down": limit_down,
        "life_high": life_high,
        "life_low": life_low,
    }


def _decode_header(values: list[Any]) -> tuple[str, int, int, str, list[Any]]:
    symbol = _text(values.pop(0))
    action_day, trading_day, update_time = values[:3]
    del values[:3]
    return symbol, action_day, trading_day, _time_text(update_time), values


def decode_gfex_qty(payload: bytes) -> dict[str, Any]:
    """Decode one ``gfex.fut_qty.*`` payload."""
    if len(payload) != GFEX_QTY_SIZE:
        raise ValueError(
            f"expected {GFEX_QTY_SIZE} bytes for GFEX quantity, got {len(payload)}"
        )
    values = list(GFEX_QTY.unpack(payload))
    symbol, action_day, trading_day, update_time, values = _decode_header(values)
    prices = values[:5]
    buy_open = values[5:10]
    buy_close = values[10:15]
    sell_open = values[15:20]
    sell_close = values[20:25]
    return {
        "symbol": symbol,
        "action_day": action_day,
        "trading_day": trading_day,
        "update_time": update_time,
        "levels": [
            {
                "price": prices[i],
                "buy_open_volume": buy_open[i],
                "buy_close_volume": buy_close[i],
                "sell_open_volume": sell_open[i],
                "sell_close_volume": sell_close[i],
            }
            for i in range(5)
        ],
    }


def decode_gfex_order(payload: bytes) -> dict[str, Any]:
    """Decode one ``gfex.fut_order.*`` payload."""
    if len(payload) != GFEX_ORDER_SIZE:
        raise ValueError(
            f"expected {GFEX_ORDER_SIZE} bytes for GFEX order, got {len(payload)}"
        )
    values = list(GFEX_ORDER.unpack(payload))
    symbol, action_day, trading_day, update_time, values = _decode_header(values)
    bid_price = values[0]
    bid_order_qty = values[1:11]
    ask_price = values[11]
    ask_order_qty = values[12:22]
    return {
        "symbol": symbol,
        "action_day": action_day,
        "trading_day": trading_day,
        "update_time": update_time,
        "bid_price": bid_price,
        "bid_order_qty": bid_order_qty,
        "ask_price": ask_price,
        "ask_order_qty": ask_order_qty,
    }


DECODERS = {
    "gfex.fut_l2.": (decode_gfex_l2, GFEX_L2_SIZE),
    "gfex.fut_qty.": (decode_gfex_qty, GFEX_QTY_SIZE),
    "gfex.fut_order.": (decode_gfex_order, GFEX_ORDER_SIZE),
}


def decoder_for(subject: str):
    for prefix, decoder_info in DECODERS.items():
        if subject.startswith(prefix):
            return decoder_info
    return None


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--url",
        default=os.getenv("NATS_URL", "nats://nats.example.com:4222"),
        help="NATS client URL (default: NATS_URL or example endpoint)",
    )
    parser.add_argument(
        "--user", default=os.getenv("NATS_USER"), help="NATS username"
    )
    parser.add_argument(
        "--password", default=os.getenv("NATS_PASSWORD"), help="NATS password"
    )
    parser.add_argument(
        "--subject",
        default=os.getenv("GFEX_SUBJECT", "gfex.fut_l2.>"),
        help="one subject; default: gfex.fut_l2.>",
    )
    parser.add_argument(
        "--all",
        action="store_true",
        help="subscribe to L2, quantity, and top-ten order subjects",
    )
    parser.add_argument(
        "--limit",
        type=int,
        default=0,
        help="stop after this many messages; 0 means run until Ctrl-C",
    )
    parser.add_argument(
        "--pending-msgs",
        type=int,
        default=100_000,
        help="client callback queue limit",
    )
    parser.add_argument(
        "--pending-bytes",
        type=int,
        default=64 * 1024 * 1024,
        help="client callback byte queue limit",
    )
    return parser.parse_args()


async def run(args: argparse.Namespace) -> None:
    nc = NATS()
    stopped = asyncio.Event()
    message_count = 0
    decode_errors = 0

    async def on_error(error: Exception) -> None:
        print(f"[NATS error] {error}", file=sys.stderr)

    async def on_disconnected() -> None:
        print("[NATS] disconnected; waiting for reconnect", file=sys.stderr)

    async def on_reconnected() -> None:
        print(f"[NATS] reconnected to {nc.connected_url.netloc}", file=sys.stderr)

    async def on_closed() -> None:
        print("[NATS] connection closed", file=sys.stderr)

    connect_kwargs: dict[str, Any] = {
        "servers": [args.url],
        "allow_reconnect": True,
        "max_reconnect_attempts": -1,
        "reconnect_time_wait": 2,
        "connect_timeout": 5,
        "name": "mdlservice-gfex-l2-python-example",
        "error_cb": on_error,
        "disconnected_cb": on_disconnected,
        "reconnected_cb": on_reconnected,
        "closed_cb": on_closed,
    }
    # Passing credentials separately avoids URL parsing issues when a password
    # contains '@', ':', or other URL-reserved characters.
    if args.user is not None:
        connect_kwargs["user"] = args.user
    if args.password is not None:
        connect_kwargs["password"] = args.password

    print(f"[NATS] connecting to {args.url}")
    await nc.connect(**connect_kwargs)

    async def on_message(msg: Any) -> None:
        nonlocal message_count, decode_errors
        decoder_info = decoder_for(msg.subject)
        if decoder_info is None:
            decode_errors += 1
            print(f"[decode error] unsupported subject={msg.subject}", file=sys.stderr)
            return
        decoder, _ = decoder_info
        try:
            record = decoder(msg.data)
        except ValueError as error:
            decode_errors += 1
            print(f"[decode error] subject={msg.subject}: {error}", file=sys.stderr)
            return

        message_count += 1
        print(json.dumps({"subject": msg.subject, **record}, ensure_ascii=False))
        if args.limit and message_count >= args.limit:
            stopped.set()

    subjects = (
        ["gfex.fut_l2.>", "gfex.fut_qty.>", "gfex.fut_order.>"]
        if args.all
        else [args.subject]
    )
    for subject in subjects:
        await nc.subscribe(
            subject,
            cb=on_message,
            pending_msgs_limit=args.pending_msgs,
            pending_bytes_limit=args.pending_bytes,
        )
    await nc.flush()
    print(
        f"[NATS] subscribed to {', '.join(subjects)}; "
        "press Ctrl-C to stop",
        file=sys.stderr,
    )

    try:
        await stopped.wait()
    finally:
        if not nc.is_closed:
            await nc.drain()
        print(
            f"[NATS] received={message_count}, decode_errors={decode_errors}",
            file=sys.stderr,
        )


def main() -> None:
    args = parse_args()
    try:
        asyncio.run(run(args))
    except KeyboardInterrupt:
        print("\n[NATS] stopped by user", file=sys.stderr)
    except Exception as error:
        print(f"[fatal] {error}", file=sys.stderr)
        raise SystemExit(1) from error


if __name__ == "__main__":
    main()
