#!/usr/bin/env python3
"""Subscribe to the stock Level 2 NATS subjects with NumPy packed dtypes."""

import argparse
import asyncio
import datetime as dt
import sys

import numpy as np
from nats.aio.client import Client as NATS


L2 = np.dtype([
    ("symbol", "S32"), ("market", "u1"), ("date", "i4"), ("time", "i4"),
    ("pre_close", "u4"), ("open", "u4"), ("high", "u4"), ("low", "u4"),
    ("last", "u4"), ("volume", "i8"), ("turnover", "i8"), ("num_trades", "i8"),
    ("ask_price", "10u4"), ("ask_vol", "10i8"), ("bid_price", "10u4"),
    ("bid_vol", "10i8"), ("ask_num_orders", "10i4"), ("bid_num_orders", "10i4"),
    ("total_ask_vol", "i8"), ("total_bid_vol", "i8"), ("avg_ask_price", "u4"),
    ("avg_bid_price", "u4"), ("limit_up", "u4"), ("limit_down", "u4"),
    ("iopv", "u4"), ("pre_close_iopv", "u4"), ("trading_phase", "S8"),
    ("is_after_hours", "i4"),
])
TRANS = np.dtype([
    ("symbol", "S32"), ("market", "u1"), ("date", "i4"), ("time", "i4"),
    ("index", "i8"), ("price", "u4"), ("volume", "i8"), ("turnover", "i8"),
    ("buy_id", "i8"), ("sell_id", "i8"), ("bs_flag", "S1"), ("order_kind", "S1"),
    ("function_code", "S1"), ("channel", "i4"),
])
ORDER = np.dtype([
    ("symbol", "S32"), ("market", "u1"), ("date", "i4"), ("time", "i4"),
    ("index", "i8"), ("order_no", "i8"), ("price", "u4"), ("volume", "i8"),
    ("bs_flag", "S1"), ("order_kind", "S1"), ("function_code", "S1"), ("channel", "i4"),
])
INDEX = np.dtype([
    ("symbol", "S32"), ("market", "u1"), ("date", "i4"), ("time", "i4"),
    ("pre_close", "u4"), ("open", "u4"), ("high", "u4"), ("low", "u4"),
    ("last", "u4"), ("volume", "i8"), ("turnover", "i8"), ("close", "u4"),
    ("num_trades", "i8"),
])
DECODERS = {
    "stk.l2.": (L2, "SNAPSHOT"), "stk.trans.": (TRANS, "TRANS"),
    "stk.order.": (ORDER, "ORDER"), "stk.index.": (INDEX, "INDEX"),
}


def time_text(value: int) -> str:
    ms = value % 1000
    value //= 1000
    seconds, value = value % 100, value // 100
    minutes, hours = value % 100, value // 100
    return f"{hours:02d}:{minutes:02d}:{seconds:02d}.{ms:03d}"


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


async def run(args: argparse.Namespace) -> None:
    nc = NATS()
    await nc.connect(
        servers=[args.url], user=args.user, password=args.password,
        allow_reconnect=True, max_reconnect_attempts=-1, reconnect_time_wait=2,
    )
    print(f"[+] connected: {args.url}")

    async def on_message(msg):
        decoder = decoder_for(msg.subject)
        if decoder is None:
            return
        dtype, label = decoder
        if len(msg.data) != dtype.itemsize:
            print(f"[-] {msg.subject}: expected {dtype.itemsize} bytes, got {len(msg.data)}")
            return
        record = np.frombuffer(msg.data, dtype=dtype)[0]
        symbol = record["symbol"].decode("ascii", errors="replace").rstrip("\x00")
        price = record["last"] / 10000.0 if "last" in dtype.names else record["price"] / 10000.0
        print(f"[{dt.datetime.now():%H:%M:%S.%f}][{label}] {msg.subject} {symbol} price={price:.4f} time={time_text(record['time'])}")

    for subject in args.subjects:
        await nc.subscribe(subject, cb=on_message, pending_msgs_limit=1_000_000, pending_bytes_limit=256 * 1024 * 1024)
        print(f"[+] subscribed: {subject}")
    try:
        while True:
            await asyncio.sleep(1)
    finally:
        await nc.drain()


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--url", default="nats://quote5.base32.cn:4222")
    parser.add_argument("--user", default="level2_test")
    parser.add_argument("--password", default="level2_test")
    parser.add_argument("--subjects", nargs="+", default=["stk.l2.600519", "stk.trans.600519", "stk.order.600519"])
    return parser.parse_args()


if __name__ == "__main__":
    try:
        asyncio.run(run(parse_args()))
    except KeyboardInterrupt:
        sys.exit(0)
