"""Read unregistered ONT serials from the authorized, pinned C-DATA OLT.

This script accepts no user CLI text.  It sends exactly the read-only command
sequence validated from the operator console for GPON 0/0, PON 1.
"""
import asyncio
import base64
import hashlib
import hmac
import os
import re
import socket
import time

from dotenv import load_dotenv
import paramiko

load_dotenv()

from app.core.database import AsyncSessionLocal
from app.core.device_credentials import decrypt_secret
from app.core.ssh_identity import _normalize_sha256_fingerprint
from app.models.models import DeviceCredential, OLTDevice

COMMANDS = ("enable", "config", "interface gpon 0/0", "show ont autofind 1 all")
SERIAL = re.compile(r"^\s*SN\s*:\s*([A-Za-z0-9_-]{8,64})(?:\s|$)", re.MULTILINE)
VENDOR = re.compile(r"^\s*Vendor ID\s*:\s*([A-Za-z0-9_-]{1,64})\s*$", re.MULTILINE)


def fingerprint(key: object) -> str:
    return base64.b64encode(hashlib.sha256(key.asbytes()).digest()).decode("ascii").rstrip("=")


def read_available(channel: paramiko.Channel, seconds: float = 1.2) -> str:
    deadline = time.monotonic() + seconds
    chunks: list[bytes] = []
    while time.monotonic() < deadline:
        if channel.recv_ready():
            chunks.append(channel.recv(65535))
            deadline = time.monotonic() + 0.25
        else:
            time.sleep(0.05)
    return b"".join(chunks).decode("utf-8", errors="replace")


def run_discovery(host: str, port: int, username: str, encrypted_secret: str) -> list[tuple[str, str | None]]:
    expected = _normalize_sha256_fingerprint(os.getenv("OLT_1_SSH_HOST_KEY_SHA256"))
    if expected is None:
        raise RuntimeError("pinned host key missing")
    sock = socket.create_connection((host, port), timeout=8)
    transport = paramiko.Transport(sock)
    channel = None
    try:
        transport.start_client(timeout=8)
        if not hmac.compare_digest(fingerprint(transport.get_remote_server_key()), expected):
            raise RuntimeError("pinned host key mismatch")
        secret = decrypt_secret(encrypted_secret)
        try:
            transport.auth_password(username, secret, fallback=False)
        finally:
            secret = ""
        if not transport.is_authenticated():
            raise RuntimeError("SSH authentication failed")
        channel = transport.open_session(timeout=8)
        channel.get_pty(term="vt100", width=160, height=40)
        channel.invoke_shell()
        read_available(channel)
        transcript = ""
        command_statuses: list[str] = []
        for command in COMMANDS:
            channel.send(command + "\r")
            response = read_available(channel, seconds=5.0 if command == "show ont autofind 1 all" else 1.2)
            transcript += response
            command_statuses.append("error" if any(marker in response.lower() for marker in ("invalid input", "incomplete command", "unknown command")) else ("response" if response.strip() else "empty"))
        serials = SERIAL.findall(transcript)
        vendors = VENDOR.findall(transcript)
        diagnostics = {
            "autofind_response": "Aging time of the automatically found ONTs" in transcript,
            "sn_label_present": bool(re.search(r"\bSN\b", transcript)),
            "cli_error": any(marker in transcript.lower() for marker in ("invalid input", "incomplete command", "unknown command")),
            "command_statuses": command_statuses,
        }
        seen: set[str] = set()
        results: list[tuple[str, str | None]] = []
        for index, serial in enumerate(serials):
            if serial not in seen:
                seen.add(serial)
                results.append((serial, vendors[index] if index < len(vendors) else None))
        return results, diagnostics
    finally:
        if channel is not None:
            channel.close()
        transport.close()


async def main() -> None:
    async with AsyncSessionLocal() as db:
        olt = await db.get(OLTDevice, 1)
        if not olt or (olt.vendor, olt.model) != ("C-DATA", "FD1601S-B1") or not olt.credential_id:
            raise RuntimeError("authorized OLT target unavailable")
        credential = await db.get(DeviceCredential, olt.credential_id)
        if not credential or not credential.is_active:
            raise RuntimeError("active OLT credential unavailable")
        rows, diagnostics = await asyncio.to_thread(run_discovery, olt.host, olt.api_port, credential.username, credential.encrypted_secret)
    print(f"auto_sn_count={len(rows)}")
    print("discovery_response=" + ("present" if diagnostics["autofind_response"] else "missing"))
    print("sn_label=" + ("present" if diagnostics["sn_label_present"] else "missing"))
    print("cli_command_error=" + ("present" if diagnostics["cli_error"] else "none"))
    print("command_statuses=" + ",".join(diagnostics["command_statuses"]))
    for serial, vendor in rows:
        print(f"auto_sn={serial}|vendor={vendor or 'unknown'}")


asyncio.run(main())
