"""RTR over SSH with explicit host trust and client authentication (asyncio).

Install rpkiparrot[ssh], or [cli,ssh] for --cli. The local peer and temporary
keys are demonstration fixtures; this library does not provide an RTR server.
For a real upstream, use its independently verified known_hosts and authorized
credentials in SshConfig and retain the Client lifecycle shown below.
"""

from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path
from tempfile import TemporaryDirectory

import anyio
import asyncssh

from rpkiparrot import Client
from rpkiparrot.config import ClientConfig, RtrSourceConfig, SourceGroupConfig, SshConfig
from rpkiparrot.models import PayloadKind
from rpkiparrot.validation import validate_origin


async def demo(cli: bool = False) -> None:
    """Authenticate a real SSH channel and accept one complete RFC 8210 update."""
    host_key = asyncssh.generate_private_key("ssh-ed25519")
    client_key = asyncssh.generate_private_key("ssh-ed25519")

    class Server(asyncssh.SSHServer):
        """Restrict the fixture to the example's single user and public key."""

        def begin_auth(self, username: str) -> bool:
            return True

        def public_key_auth_supported(self) -> bool:
            return True

        def validate_public_key(self, username: str, key: asyncssh.SSHKey) -> bool:
            return username == "rtr-reader" and key == client_key.convert_to_public()

    async def peer(process: asyncssh.SSHServerProcess[bytes]) -> None:
        # Independent bytes transcribed from RFC 8210 §§5.2-5.4,5.7.
        try:
            if process.subsystem != "rpki-rtr" or process.command is not None:
                raise RuntimeError("expected the rpki-rtr subsystem without exec")
            if await process.stdin.readexactly(8) != bytes.fromhex("0102000000000008"):
                raise RuntimeError("expected a v1 Reset Query")
            process.stdout.write(
                bytes.fromhex(
                    "0103002a00000008"
                    "010400000000001401181800c00002000000fbf0"
                    "0107002a00000018000000070000001e0000000100000258"
                )
            )
            await process.stdout.drain()
            await process.stdin.read()
        finally:
            process.exit(0)

    with TemporaryDirectory(prefix="rpkiparrot-ssh-example-") as directory:
        root = Path(directory)
        private = root / "client-key"
        private.write_bytes(client_key.export_private_key())
        private.chmod(0o600)
        async with await asyncssh.create_server(
            Server,
            "127.0.0.1",
            0,
            server_host_keys=[host_key],
            process_factory=peer,
            encoding=None,
        ) as server:
            port = server.get_port()
            known = root / "known_hosts"
            known.write_bytes(f"[127.0.0.1]:{port} ".encode() + host_key.export_public_key())
            if cli:
                with anyio.fail_after(20):
                    completed = await anyio.run_process(
                        [
                            sys.executable,
                            "-I",
                            "-m",
                            "rpkiparrot",
                            "--format",
                            "json",
                            "--rtr",
                            f"127.0.0.1:{port}",
                            "--ssh",
                            "--ssh-user",
                            "rtr-reader",
                            "--ssh-known-hosts",
                            str(known),
                            "--ssh-key",
                            str(private),
                            "--trust-profile-id",
                            "example-host-key-v1",
                            "validate",
                            "origin",
                            "192.0.2.0/24",
                            "64496",
                        ]
                    )
                result = json.loads(completed.stdout)
                assert result["data"]["status"] == "valid" and result["error"] is None
                print(completed.stdout.decode(), end="")
            else:
                source = RtrSourceConfig(
                    id="ssh-cache",
                    host="127.0.0.1",
                    port=port,
                    transport="ssh",
                    trust_profile_id="example-host-key-v1",
                    ssh=SshConfig(username="rtr-reader", known_hosts=known, client_keys=(private,)),
                )
                config = ClientConfig(
                    groups=(
                        SourceGroupConfig(
                            id="primary",
                            priority=0,
                            sources=(source,),
                        ),
                    )
                )
                async with anyio.create_task_group() as tasks, Client(config) as client:
                    await tasks.start(client.run)
                    snapshot = await client.wait_ready(
                        required=frozenset({PayloadKind.VRP}), timeout=10
                    )
                    result_value = validate_origin(snapshot, "192.0.2.0/24", 64496)
                    assert result_value.status.value == "valid"
                    print(json.dumps({"transport": "ssh", "origin": result_value.status.value}))


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--cli", action="store_true", help="exercise the installed CLI")
    anyio.run(demo, parser.parse_args().cli, backend="asyncio")
