"""Synchronize and publicly restore one RFC 8210 source on asyncio or Trio.

The tiny peer is an example/test fixture, not a production RTR server. It sends
specified wire bytes directly; the application owns both tasks and all sockets.
"""

from __future__ import annotations

import argparse
import json
import struct
from contextlib import suppress

import anyio
from anyio.abc import ByteStream, SocketAttribute

from rpkiparrot import MemoryStore
from rpkiparrot.config import FreshnessPolicy, RtrSourceConfig
from rpkiparrot.models import SourceInfo, SourceUpdate
from rpkiparrot.rtr.client import RtrSession
from rpkiparrot.validation import validate_origin


class Sink:
    """Connect the low-level session to an application-owned memory store."""

    def __init__(self) -> None:
        self.store = MemoryStore()
        self.ready = anyio.Event()
        self.accepted: SourceUpdate | None = None

    async def publish(self, update: SourceUpdate) -> None:
        """Publish complete original records without restarting their lifetime."""
        self.store.replace_source(
            update.source.id,
            update.dataset,
            freshness=FreshnessPolicy(valid_until=update.source.expires_at),
        )
        # No await separates the complete store commit, acknowledgement value
        # and return. Persist this original value in an application if needed.
        self.accepted = update
        self.ready.set()

    async def report(self, status: SourceInfo) -> None:
        """Remove revoked data; retrying connectivity alone retains valid data."""
        if any(value.value in ("unavailable", "expired") for value in status.capabilities.values()):
            self.store.remove_source(status.id)
            self.accepted = None


async def peer(stream: ByteStream) -> None:
    """Accept Reset or restored Serial Query using independent RFC wire bytes."""
    async with stream:
        request = bytearray()
        while len(request) < 8:
            request.extend(await stream.receive(8 - len(request)))
        if request[1] == 1:
            while len(request) < 12:
                request.extend(await stream.receive(12 - len(request)))
            if struct.unpack("!BBHII", request) != (1, 1, 42, 12, 7):
                raise RuntimeError("example cache received incorrect restored session/serial")
        elif request != bytes.fromhex("0102000000000008"):
            raise RuntimeError("example cache expected an RTR v1 query")
        wire = struct.pack("!BBHI", 1, 3, 42, 8)
        if request[1] == 2:
            wire += struct.pack("!BBHIBBBBII", 1, 4, 0, 20, 1, 24, 24, 0, 0xC0000200, 64496)
        wire += struct.pack("!BBHIIIII", 1, 7, 42, 24, 7, 30, 1, 600)
        await stream.send(wire)
        with suppress(anyio.EndOfStream):
            await stream.receive()


async def main() -> None:
    """Own the listener and session in one managed task group."""
    listener = await anyio.create_tcp_listener(local_host="127.0.0.1", local_port=0)
    port = listener.extra(SocketAttribute.local_port)
    restored: SourceUpdate | None = None
    async with listener, anyio.create_task_group() as tasks:
        tasks.start_soon(listener.serve, peer)
        config = RtrSourceConfig(id="demo", host="127.0.0.1", port=port)
        for _ in range(2):
            sink = Sink()
            async with RtrSession(config, restored=restored) as session:
                await tasks.start(session.run, sink)
                with anyio.fail_after(5):
                    await sink.ready.wait()
                restored = sink.accepted
                assert restored is not None
                result = validate_origin(sink.store.snapshot(), "192.0.2.0/24", 64496)
        assert restored is not None and restored.reason == "incremental"
        print(
            json.dumps(
                {"origin": result.status.value, "source": config.id, "restored": restored.reason}
            )
        )
        tasks.cancel_scope.cancel()


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--backend", choices=("asyncio", "trio"), default="asyncio")
    anyio.run(main, backend=parser.parse_args().backend)
