"""Real localhost CoAP GET, PUT and Observe traffic through a UDP byte tap."""
import argparse
import asyncio

import aiocoap
from aiocoap import resource

PROXY_PORT = 17680
SERVER_PORT = 17681


class ByteTap(asyncio.DatagramProtocol):
    def __init__(self):
        self.transport = None
        self.client = None
        self.frames = []

    def connection_made(self, transport):
        self.transport = transport

    def datagram_received(self, data, address):
        if address[1] == SERVER_PORT:
            self.frames.append(("S>C", data))
            if self.client:
                self.transport.sendto(data, self.client)
        else:
            self.client = address
            self.frames.append(("C>S", data))
            self.transport.sendto(data, ("127.0.0.1", SERVER_PORT))


def describe(frame):
    direction, data = frame
    token_len = data[0] & 15
    code = f"{data[1] >> 5}.{data[1] & 31:02d}"
    return (f"UDP {direction} {len(data)} bytes hex={data.hex()}\n"
            f"CoAP ver={data[0] >> 6} type={(data[0] >> 4) & 3} "
            f"code={code} mid=0x{int.from_bytes(data[2:4], 'big'):04x} "
            f"token={data[4:4 + token_len].hex()}")


class Discovery(resource.Resource):
    async def render_get(self, request):
        response = aiocoap.Message(payload=b'</temperature>;obs,</target>;rt="setpoint"')
        response.opt.content_format = 40
        return response


class Temperature(resource.ObservableResource):
    def __init__(self):
        super().__init__()
        self.celsius = 21.5

    async def render_get(self, request):
        response = aiocoap.Message(payload=f"{self.celsius:.1f}".encode())
        response.opt.content_format = 0
        return response


class Target(resource.Resource):
    def __init__(self):
        super().__init__()
        self.celsius = 20.0

    async def render_get(self, request):
        return aiocoap.Message(payload=f"{self.celsius:.1f}".encode())

    async def render_put(self, request):
        try:
            value = float(request.payload.decode())
            if not 5 <= value <= 35:
                raise ValueError()
        except ValueError:
            return aiocoap.Message(code=aiocoap.BAD_REQUEST, payload=b"target must be 5..35 C")
        self.celsius = value
        return aiocoap.Message(code=aiocoap.CHANGED, payload=f"{value:.1f}".encode())


async def run(step):
    site = resource.Site()
    temperature = Temperature()
    site.add_resource(("sensors",), Discovery())
    site.add_resource(("temperature",), temperature)
    site.add_resource(("target",), Target())
    server = await aiocoap.Context.create_server_context(site, bind=("127.0.0.1", SERVER_PORT))
    loop = asyncio.get_running_loop()
    transport, tap = await loop.create_datagram_endpoint(ByteTap, local_addr=("127.0.0.1", PROXY_PORT))
    client = await aiocoap.Context.create_client_context()
    base = f"coap://127.0.0.1:{PROXY_PORT}"
    outputs = {}
    try:
        for number, method, path, payload in [
            (1, aiocoap.GET, "sensors", b""),
            (2, aiocoap.GET, "temperature", b""),
            (3, aiocoap.PUT, "target", b"22.0"),
            (4, aiocoap.GET, "target", b""),
        ]:
            start = len(tap.frames)
            request = aiocoap.Message(code=method, uri=f"{base}/{path}", payload=payload)
            response = await asyncio.wait_for(client.request(request).response, 3)
            await asyncio.sleep(.04)
            frames = tap.frames[start:]
            outputs[number] = [
                f"STEP {number}: {method} /{path}",
                *(describe(frame) for frame in frames[:2]),
                f"application response={response.code} payload={response.payload.decode()}",
                "Bytes above came from this run's UDP datagrams, not a packet example.",
            ]
        start = len(tap.frames)
        request = aiocoap.Message(code=aiocoap.GET, uri=f"{base}/temperature", observe=0)
        requester = client.request(request)
        first = await asyncio.wait_for(requester.response, 3)
        notifications = []
        pending = asyncio.Queue()
        requester.observation.register_callback(pending.put_nowait)
        for value in (22.0, 22.5):
            temperature.celsius = value
            temperature.updated_state()
            notifications.append(await asyncio.wait_for(pending.get(), 3))
        await asyncio.sleep(.04)
        observe_frames = tap.frames[start:]
        requester.observation.cancel()
        outputs[5] = [
            "STEP 5: GET /temperature Observe=0, then two server updates",
            *(describe(frame) for frame in observe_frames[:2]),
            f"initial response={first.payload.decode()} observe={first.opt.observe}",
            f"notification 1={notifications[0].payload.decode()} observe={notifications[0].opt.observe}",
            f"notification 2={notifications[1].payload.decode()} observe={notifications[1].opt.observe}",
            "Notifications above came from this run's CoAP server over local UDP.",
        ]
        response_frames = [frame for frame in observe_frames if frame[0] == "S>C"]
        tokens = [frame[1][4:4 + (frame[1][0] & 15)].hex() for frame in response_frames]
        outputs[6] = [
            "STEP 6: correlate the Observe token and UDP frames",
            f"server datagrams={len(response_frames)}",
            *(f"response {n}: token={token} code={response_frames[n - 1][1][1] >> 5}.{response_frames[n - 1][1][1] & 31:02d} hex={response_frames[n - 1][1].hex()}" for n, token in enumerate(tokens, 1)),
            f"same token across initial response and updates={len(set(tokens)) == 1}",
            "Hex rows are original aiocoap UDP datagrams relayed by the byte tap.",
        ]
        for number in ([step] if step else sorted(outputs)):
            print("\n".join(outputs[number]), flush=True)
    finally:
        transport.close()
        await client.shutdown()
        await server.shutdown()


if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument("--step", type=int, choices=range(1, 7))
    args = parser.parse_args()
    asyncio.run(run(args.step))
