import logging
import os
import tempfile

from jnpr.junos import Device as PyEZDevice
from jnpr.junos.utils.config import Config as PyEZConfig
from jnpr.junos.exception import (
    ConnectAuthError,
    ConnectTimeoutError,
    ConnectRefusedError,
    ConnectError,
    LockError,
    CommitError,
    ConfigLoadError,
)

from . import crypto

logger = logging.getLogger(__name__)


class JunosClientError(Exception):
    """Base exception for all Junos integration errors."""


class JunosAuthError(JunosClientError):
    pass


class JunosConnectError(JunosClientError):
    pass


class JunosLockedError(JunosClientError):
    pass


class JunosConfigError(JunosClientError):
    """The submitted set/delete statements were rejected while loading."""


class JunosDiffMismatchError(JunosClientError):
    """The device changed between preview and commit; caller must re-preview."""


class JunosCommitCheckError(JunosClientError):
    """commit-check failed; nothing was applied to the device."""


class JunosCommitError(JunosClientError):
    pass


def _xtext(el):
    if el is None or el.text is None:
        return None
    return el.text.strip()


class JunosSession:
    """Short-lived wrapper around a single PyEZ connection to one device.

    Always used as a context manager so the underlying SSH/NETCONF
    connection (and any temporary SSH key file) is cleaned up.
    """

    def __init__(self, device):
        self.device_record = device
        self.dev = None
        self._temp_key_path = None

    def _credentials(self):
        secret = crypto.decrypt(self.device_record.secret_encrypted)
        if self.device_record.auth_type == "sshkey":
            fd, path = tempfile.mkstemp(prefix="junos_key_")
            with os.fdopen(fd, "w") as f:
                f.write(secret)
            self._temp_key_path = path
            passphrase = None
            if self.device_record.key_passphrase_encrypted:
                passphrase = crypto.decrypt(self.device_record.key_passphrase_encrypted)
            return {"ssh_private_key_file": path, "password": passphrase}
        return {"password": secret}

    def __enter__(self):
        creds = self._credentials()
        try:
            self.dev = PyEZDevice(
                host=self.device_record.hostname,
                port=self.device_record.port,
                user=self.device_record.username,
                gather_facts=False,
                auto_probe=5,
                **creds,
            )
            self.dev.open()
        except ConnectAuthError as exc:
            self._cleanup_key_file()
            raise JunosAuthError(
                f"Authentication failed for {self.device_record.hostname}"
            ) from exc
        except (ConnectTimeoutError, ConnectRefusedError) as exc:
            self._cleanup_key_file()
            raise JunosConnectError(
                f"Could not reach {self.device_record.hostname}:{self.device_record.port}"
            ) from exc
        except ConnectError as exc:
            self._cleanup_key_file()
            raise JunosConnectError(str(exc)) from exc
        return self

    def __exit__(self, exc_type, exc_val, exc_tb):
        if self.dev is not None:
            try:
                self.dev.close()
            except Exception:
                logger.exception(
                    "Error closing Junos session to %s", self.device_record.hostname
                )
        self._cleanup_key_file()
        return False

    def _cleanup_key_file(self):
        if self._temp_key_path and os.path.exists(self._temp_key_path):
            try:
                os.remove(self._temp_key_path)
            except OSError:
                logger.exception("Could not remove temp key file %s", self._temp_key_path)
            self._temp_key_path = None

    def get_facts(self) -> dict:
        self.dev.facts_refresh()
        facts = self.dev.facts
        uptime = None
        re0 = facts.get("RE0") or {}
        if isinstance(re0, dict):
            uptime = re0.get("up_time")
        return {
            "hostname": facts.get("hostname"),
            "model": facts.get("model"),
            "version": facts.get("version"),
            "serialnumber": facts.get("serialnumber"),
            "uptime": uptime,
        }

    def get_interface_summary(self) -> list:
        result = self.dev.rpc.get_interface_information(terse=True)
        interfaces = []
        for phy in result.findall(".//physical-interface"):
            name = _xtext(phy.find("name"))
            if not name:
                continue
            interfaces.append(
                {
                    "name": name,
                    "admin_status": _xtext(phy.find("admin-status")),
                    "oper_status": _xtext(phy.find("oper-status")),
                    "description": _xtext(phy.find("description")),
                }
            )
        return interfaces

    def get_config(self, fmt="text") -> str:
        result = self.dev.rpc.get_config(options={"format": fmt})
        return result.text if hasattr(result, "text") else str(result)

    def _locked_config(self):
        cu = PyEZConfig(self.dev)
        try:
            cu.lock()
        except LockError as exc:
            raise JunosLockedError(
                "Configuration is locked by another session"
            ) from exc
        return cu

    def _unlock(self, cu):
        try:
            cu.unlock()
        except Exception:
            logger.exception("Error unlocking config on %s", self.device_record.hostname)

    def preview_diff(self, config_lines: str, config_format="set") -> str:
        """Lock, load the candidate statements, return the diff, then unlock."""
        cu = self._locked_config()
        try:
            try:
                cu.load(config_lines, format=config_format)
            except ConfigLoadError as exc:
                raise JunosConfigError(str(exc)) from exc
            return cu.diff() or ""
        finally:
            self._unlock(cu)

    def commit(
        self,
        config_lines: str,
        comment: str,
        expected_diff=None,
        config_format="set",
        confirm_minutes=None,
    ) -> str:
        """Re-lock, re-load the same statements, verify the diff still
        matches what the user previewed, run commit-check, then commit.

        If confirm_minutes is set, this is a "commit confirmed": Junos will
        automatically revert to the prior configuration after that many
        minutes unless confirm_pending_commit() is called first.
        """
        cu = self._locked_config()
        try:
            try:
                cu.load(config_lines, format=config_format)
            except ConfigLoadError as exc:
                raise JunosConfigError(str(exc)) from exc

            diff = cu.diff() or ""
            if expected_diff is not None and diff.strip() != expected_diff.strip():
                raise JunosDiffMismatchError(
                    "The device configuration changed since you previewed this "
                    "change. Please preview again before committing."
                )

            try:
                cu.commit_check()
            except CommitError as exc:
                raise JunosCommitCheckError(str(exc)) from exc

            try:
                commit_kwargs = {"comment": comment}
                if confirm_minutes:
                    commit_kwargs["confirm"] = confirm_minutes
                cu.commit(**commit_kwargs)
            except CommitError as exc:
                raise JunosCommitError(str(exc)) from exc

            return diff
        finally:
            self._unlock(cu)

    def confirm_pending_commit(self, comment: str) -> None:
        """Confirm a prior 'commit confirmed' so it becomes permanent.

        This issues a plain commit with no candidate changes loaded, which
        is exactly how Junos CLI's bare `commit` confirms a pending
        confirmed commit.
        """
        cu = self._locked_config()
        try:
            try:
                cu.commit(comment=comment)
            except CommitError as exc:
                raise JunosCommitError(str(exc)) from exc
        finally:
            self._unlock(cu)

    def rollback(self, rollback_id: int, comment: str) -> str:
        cu = self._locked_config()
        try:
            cu.rollback(rollback_id)
            diff = cu.diff() or ""
            try:
                cu.commit(comment=comment)
            except CommitError as exc:
                raise JunosCommitError(str(exc)) from exc
            return diff
        finally:
            self._unlock(cu)
