"""Read-only grammar probe for occupied C-DATA ONU IDs on PON 1."""
import asyncio
import os
import re

from dotenv import load_dotenv
load_dotenv()

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


ID_PATTERNS = (
    re.compile(r"^\s*(?:ONT|ONU)\s*ID\s*:\s*(\d+)\s*$", re.MULTILINE | re.IGNORECASE),
    re.compile(r"^\s*(?:ONT|ONU)\s+(\d+)\s*[: ]", re.MULTILINE | re.IGNORECASE),
)


def read_inventory(host, port, username, encrypted_secret, expected):
    import hmac
    import paramiko
    import socket
    expected = _normalize_sha256_fingerprint(expected)
    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("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("authentication rejected")
        channel = transport.open_session(timeout=8)
        channel.get_pty(term="vt100", width=160, height=40)
        channel.invoke_shell(); _read(channel, 1.2)
        transcript = ""
        for command, wait in (("enable", 1.2), ("config", 1.2), ("interface gpon 0/0", 1.2), ("show ont info 1 all", 6.0)):
            channel.send(command + "\r"); transcript += _read(channel, wait)
        if any(m in transcript.lower() for m in ("invalid input", "incomplete command", "unknown command")):
            return "grammar_rejected", [], None, [], "cli_grammar_rejected"
        total_match = re.search(r"^\s*Total\s*:\s*(\d+)\s*$", transcript, re.MULTILINE | re.IGNORECASE)
        total = int(total_match.group(1)) if total_match else None
        found = set()
        for pattern in ID_PATTERNS:
            found.update(int(value) for value in pattern.findall(transcript) if 1 <= int(value) <= 128)
        labels = []
        for line in transcript.splitlines():
            if ":" in line:
                label = line.split(":", 1)[0].strip()
                if label and len(label) <= 80:
                    labels.append(label)
        error_match = re.search(r"^\s*Error\s*:\s*([^\r\n]{1,120})", transcript, re.MULTILINE | re.IGNORECASE)
        error_kind = "none" if error_match is None else re.sub(r"[^A-Za-z0-9 _.-]", "", error_match.group(1)).strip()[:120]
        return "response_received", sorted(found), total, labels[:40], error_kind
    finally:
        if channel is not None: channel.close()
        transport.close()


async def main():
    async with AsyncSessionLocal() as db:
        olt = await db.get(OLTDevice, 1)
        cred = await db.get(DeviceCredential, olt.credential_id)
        state, occupied, declared_total, labels, error_kind = await asyncio.to_thread(read_inventory, olt.host, olt.api_port, cred.username, cred.encrypted_secret, os.getenv("OLT_1_SSH_HOST_KEY_SHA256"))
    print(f"inventory_state={state}")
    print("inventory_error=" + error_kind)
    print("inventory_labels=" + ",".join(labels))
    print(f"declared_total={'unknown' if declared_total is None else declared_total}")
    print(f"occupied_ids_count={len(occupied)}")
    print("lowest_available=" + str(next((item for item in range(1,129) if item not in set(occupied)), 0)))


asyncio.run(main())
