#!/usr/bin/env python3
"""Linkalyst LAN UDP echo peer. Python >= 3.9, standard library only.

Default: loopback only. Bind an explicit private LAN address for iPhone tests.
No amplification: only valid unicast payloads carrying the access code are echoed unchanged.
This is a diagnostic fixture, NOT an encrypted or Internet-facing service.
"""
from __future__ import annotations
import argparse
import asyncio
from collections import OrderedDict
import hmac
import ipaddress
import json
from pathlib import Path
import random
import secrets
import socket
import struct
import sys
import time


def is_lan(address: str) -> bool:
    try:
        ip = ipaddress.ip_address(address.split('%', 1)[0])
        return (ip.is_private or ip.is_loopback or ip.is_link_local) and not (
            ip.is_unspecified or ip.is_multicast or ip.is_reserved and not ip.is_loopback
        )
    except ValueError:
        return False


class EchoPeer(asyncio.DatagramProtocol):
    def __init__(self, args: argparse.Namespace) -> None:
        self.args = args
        self.token = args.token.encode('ascii')
        self.transport = None
        self.rng = random.Random(args.seed)
        self.clients = OrderedDict()
        self.global_second = 0
        self.global_count = 0
        self.delayed = 0
        self.accepted = 0
        self.rejected = 0
        self.injected_drops = 0
        self.rate_drops = 0
        self.closed = False

    def connection_made(self, transport) -> None:
        self.transport = transport

    def allowed_rate(self, address) -> bool:
        # Per-address (not source-port) and global limits. Bound all bookkeeping.
        second = int(time.monotonic())
        if self.global_second != second:
            self.global_second, self.global_count = second, 0
        key = address[0]
        old_second, count = self.clients.pop(key, (second, 0))
        if old_second != second:
            count = 0
        count += 1
        self.clients[key] = (second, count)
        while len(self.clients) > 256:
            self.clients.popitem(last=False)
        self.global_count += 1
        return count <= 200 and self.global_count <= 1000

    def datagram_received(self, data: bytes, address) -> None:
        if not self.allowed_rate(address):
            self.rate_drops += 1
            return
        if (not is_lan(address[0]) or not 32 <= len(data) <= 1200
                or data[:4] != b'NSP1' or not hmac.compare_digest(data[4:20], self.token)
                or any(value != 0xA5 for value in data[32:])):
            self.rejected += 1
            return
        self.accepted += 1
        sequence = struct.unpack('!I', data[28:32])[0]
        if self.args.drop_every and (sequence + 1) % self.args.drop_every == 0:
            self.injected_drops += 1
            return
        delay = max(0, self.args.delay_ms + self.rng.uniform(-self.args.jitter_ms, self.args.jitter_ms))
        delay += self.args.delay_sequences.get(sequence, 0)
        self.deliver(data, address, delay / 1000)
        if self.args.duplicate_every and (sequence + 1) % self.args.duplicate_every == 0:
            self.deliver(data, address, delay / 1000 + 0.01)

    def deliver(self, data: bytes, address, delay: float) -> None:
        if self.delayed >= 2048:
            self.rate_drops += 1
            return
        if delay <= 0:
            if not self.closed:
                self.transport.sendto(data, address)
            return
        self.delayed += 1
        def echo() -> None:
            self.delayed -= 1
            if not self.closed:
                self.transport.sendto(data, address)
        asyncio.get_running_loop().call_later(delay, echo)

    def error_received(self, exc: Exception) -> None:
        if not self.args.quiet:
            print(f'UDP transport error: {exc}', file=sys.stderr, flush=True)

    def connection_lost(self, exc) -> None:
        self.closed = True


def parse_args(argv=None) -> argparse.Namespace:
    p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    p.add_argument('--bind', default='127.0.0.1', help='Explicit private LAN IP, default loopback only')
    p.add_argument('--port', type=int, default=9999)
    p.add_argument('--token', default=None, help='access code: 16 ASCII letters/digits; random by default')
    p.add_argument('--drop-every', type=int, default=0, help='Fault injection, zero disables')
    p.add_argument('--duplicate-every', type=int, default=0, help='Fault injection, zero disables')
    p.add_argument('--delay-ms', type=float, default=0)
    p.add_argument('--jitter-ms', type=float, default=0)
    p.add_argument('--delay-seq', action='append', default=[], metavar='SEQ:MS')
    p.add_argument('--seed', type=int, default=1)
    p.add_argument('--ready-file', type=Path, help='For integration tests; writes bound port (not the access code)')
    p.add_argument('--quiet', action='store_true')
    args = p.parse_args(argv)
    if not is_lan(args.bind):
        p.error('--bind must be a numeric private/loopback/link-local address; wildcard/public binds are disabled')
    if not 0 <= args.port <= 65535:
        p.error('port must be 0..65535 (0 picks a temporary port)')
    args.token = args.token or secrets.token_hex(8)
    if len(args.token) != 16 or not args.token.isascii() or not args.token.isalnum():
        p.error('--token (the access code) must contain exactly 16 ASCII letters/digits')
    if args.drop_every < 0 or args.duplicate_every < 0:
        p.error('fault injection periods cannot be negative')
    if not (0 <= args.delay_ms <= 30000 and 0 <= args.jitter_ms <= 30000):
        p.error('delay/jitter must be finite and within 0..30000 ms')
    args.delay_sequences = {}
    for spec in args.delay_seq:
        try:
            seq, ms = spec.split(':')
            seq, ms = int(seq), float(ms)
            if not (0 <= seq <= 0xFFFFFFFF and 0 <= ms <= 30000):
                raise ValueError()
            args.delay_sequences[seq] = ms
        except ValueError:
            p.error('--delay-seq requires SEQ:MS, e.g. 0:350')
    return args


async def main(args: argparse.Namespace) -> None:
    family = socket.AF_INET6 if ':' in args.bind else socket.AF_INET
    loop = asyncio.get_running_loop()
    transport, protocol = await loop.create_datagram_endpoint(
        lambda: EchoPeer(args), local_addr=(args.bind, args.port), family=family)
    bound = transport.get_extra_info('sockname')
    if args.ready_file:
        args.ready_file.write_text(json.dumps({'host': bound[0], 'port': bound[1]}), encoding='utf-8')
    if not args.quiet:
        print(f'Linkalyst UDP echo peer: {bound[0]}:{bound[1]}\nAccess code: {args.token}', flush=True)
        print('Private networks only. Do not forward this port from the Internet; the access code is not encryption. Press Ctrl+C to stop.', flush=True)
        if args.bind in ('127.0.0.1', '::1'):
            print('Reachable from this computer only. To test from an iPhone or iPad, restart with --bind and this computer\'s LAN IP address.', flush=True)
        if any((args.drop_every, args.duplicate_every, args.delay_ms, args.jitter_ms, args.delay_sequences)):
            print('Warning: fault injection is on; results are not a baseline of natural network quality.', flush=True)
    try:
        await asyncio.Future()
    finally:
        transport.close()
        if args.ready_file:
            args.ready_file.unlink(missing_ok=True)
        if not args.quiet:
            print(f'Accepted={protocol.accepted}; invalid={protocol.rejected}; injected drops={protocol.injected_drops}; rate/capacity drops={protocol.rate_drops}', flush=True)


if __name__ == '__main__':
    try:
        asyncio.run(main(parse_args()))
    except KeyboardInterrupt:
        pass
    except OSError as exc:
        print(f'Could not start the peer: {exc}. Check that the address belongs to this computer, that the port is free, and the firewall settings.', file=sys.stderr)
        sys.exit(1)
