# Ultralytics 🚀 AGPL-3.0 License - https://ultralytics.com/license

import json
import logging
import plistlib
import shutil
import subprocess
import sys
import threading
import time
from datetime import datetime
from pathlib import Path

from ultralytics.utils import LINUX, LOGGER, MACOS, RANK, WINDOWS


class ConsoleLogger:
    """Console output capture with batched streaming to file, API, or custom callback.

    Captures stdout/stderr and Ultralytics logger output and streams it with deduplication and configurable batching.

    Attributes:
        destination (str | Path | None): Target destination for streaming (URL, Path, or None for callback-only).
        is_api (bool): Whether the destination is an HTTP(S) API endpoint.
        batch_size (int): Number of lines to batch before flushing (default: 1 for immediate).
        flush_interval (float): Seconds between automatic flushes (default: 5.0).
        on_flush (callable | None): Optional callback function called with batched content on flush.
        active (bool): Whether console capture is currently active.

    Examples:
        File logging (immediate):
        >>> logger = ConsoleLogger("training.log")
        >>> logger.start_capture()
        >>> print("This will be logged")
        >>> logger.stop_capture()

        API streaming with batching:
        >>> logger = ConsoleLogger("https://api.example.com/logs", batch_size=10)
        >>> logger.start_capture()

        Custom callback with batching:
        >>> def my_handler(content, line_count, chunk_id):
        ...     print(f"Received {line_count} lines")
        >>> logger = ConsoleLogger(on_flush=my_handler, batch_size=5)
        >>> logger.start_capture()
    """

    def __init__(self, destination=None, batch_size=1, flush_interval=5.0, on_flush=None):
        """Initialize console logger with optional batching.

        Args:
            destination (str | Path | None): API endpoint URL (http/https), local file path, or None.
            batch_size (int): Lines to accumulate before flush (1 = immediate, higher = batched).
            flush_interval (float): Max seconds between flushes when batching.
            on_flush (callable | None): Callback(content: str, line_count: int, chunk_id: int) for custom handling.
        """
        if isinstance(destination, str) and destination.startswith("http://"):
            LOGGER.warning("ConsoleLogger destination uses plaintext HTTP; captured logs are sent unencrypted.")
        self.destination = destination
        self.is_api = isinstance(destination, str) and destination.startswith(("http://", "https://"))
        if destination is not None and not self.is_api:
            self.destination = Path(destination)

        # Batching configuration
        self.batch_size = max(1, batch_size)
        self.flush_interval = flush_interval
        self.on_flush = on_flush

        # Console capture state
        self.original_stdout = sys.stdout
        self.original_stderr = sys.stderr
        self.active = False
        self._log_handler = None  # Track handler for cleanup

        # Buffer for batching
        self.buffer = []
        self.buffer_lock = threading.Lock()
        self.flush_thread = None
        self.chunk_id = 0

        # Deduplication state
        self.last_line = ""
        self.last_time = 0.0
        self.progress = {}  # live progress bar frames by bar id, held as state instead of buffered as log lines
        self.progress_sent = {}

    def start_capture(self):
        """Start capturing console output and redirect stdout/stderr.

        Notes:
            In DDP training, only activates on rank 0/-1 to prevent duplicate logging.
        """
        if self.active or RANK not in {-1, 0}:
            return

        self.active = True
        self.stdout_capture = sys.stdout = self._ConsoleCapture(self.original_stdout, self._queue_log)
        self.stderr_capture = sys.stderr = self._ConsoleCapture(self.original_stderr, self._queue_log)

        # Hook Ultralytics logger
        try:
            self._log_handler = self._LogHandler(self._queue_log)
            logging.getLogger("ultralytics").addHandler(self._log_handler)
        except Exception:
            pass

        # Background flush thread: carries live progress frames in every mode, batched lines when batching
        if not (self.flush_thread and self.flush_thread.is_alive()):  # a worker still sleeping resumes on its own
            self.flush_thread = threading.Thread(target=self._flush_worker, daemon=True)
            self.flush_thread.start()

    def stop_capture(self):
        """Stop capturing console output and flush remaining buffer."""
        if not self.active:
            return

        self.stdout_capture.callback = self.stderr_capture.callback = None
        self.active = False
        sys.stdout = self.original_stdout
        sys.stderr = self.original_stderr

        # Remove logging handler to prevent memory leak
        if self._log_handler:
            try:
                logging.getLogger("ultralytics").removeHandler(self._log_handler)
            except Exception:
                pass
            self._log_handler = None

        # Final flush, without the frames of any bar that outlived the capture
        self.progress.clear()
        self._flush_buffer()

    def _queue_log(self, text, bar_id):
        """Queue console text with deduplication and timestamp processing, holding bar frames as state."""
        if not self.active:
            return

        current_time = time.time()

        # Handle carriage returns and strip ANSI clear-line codes (TQDM writes "\r\033[K<line>" interactively)
        if "\r" in text:
            text = text.split("\r")[-1]
        text = text.replace("\x1b[K", "")

        if bar_id is not None:  # a redraw is bar state, so only the frame a bar closed with becomes a log line
            with self.buffer_lock:
                if text:
                    self.progress[str(bar_id)] = text.rstrip()
                    return
                text = self.progress.pop(str(bar_id), "")

        lines = text.split("\n")
        if lines and lines[-1] == "":
            lines.pop()

        for line in lines:
            if not (line := line.rstrip()):
                continue  # a bare newline carries nothing once timestamped

            # General deduplication
            if line == self.last_line and current_time - self.last_time < 0.1:
                continue

            self.last_line = line
            self.last_time = current_time

            # Add timestamp if needed
            if not line.startswith("[20"):
                timestamp = datetime.now().astimezone().strftime("%Y-%m-%d %H:%M:%S")
                line = f"[{timestamp}] {line}"

            # Add to buffer and check if flush needed
            should_flush = False
            with self.buffer_lock:
                self.buffer.append(line)
                if len(self.buffer) >= self.batch_size:
                    should_flush = True

            # Flush outside lock to avoid deadlock
            if should_flush:
                self._flush_buffer()

    def _flush_worker(self):
        """Background worker that flushes buffer periodically."""
        while self.active:
            time.sleep(self.flush_interval)
            if self.active:
                self._flush_buffer()

    def _flush_buffer(self):
        """Flush buffered lines and current progress bar frames to destination and/or callback."""
        with self.buffer_lock:
            if not self.buffer and self.progress == self.progress_sent:
                return  # nothing new: an idle bar must not flush a chunk of its own
            lines = self.buffer.copy()
            self.buffer.clear()
            self.progress_sent = dict(self.progress)
            self.chunk_id += 1
            chunk_id = self.chunk_id  # Capture under lock to avoid race

        content = "\n".join(lines)
        line_count = len(lines)

        # Call custom callback if provided
        if self.on_flush:
            try:
                self.on_flush(content, line_count, chunk_id)
            except Exception:
                pass  # Silently ignore callback errors to avoid flooding stderr

        # Write to destination (file or API)
        if self.destination is not None and lines:
            self._write_destination(content)

    def _write_destination(self, content):
        """Write content to file or API destination."""
        try:
            if self.is_api:
                import requests

                payload = {"timestamp": datetime.now().astimezone().isoformat(), "message": content}
                requests.post(str(self.destination), json=payload, timeout=5)
            else:
                self.destination.parent.mkdir(parents=True, exist_ok=True)
                with self.destination.open("a", encoding="utf-8") as f:
                    f.write(content + "\n")
        except Exception as e:
            print(f"Console logger write error: {e}", file=self.original_stderr)

    class _ConsoleCapture:
        """Lightweight stdout/stderr capture."""

        __slots__ = ("callback", "original")

        def __init__(self, original, callback):
            """Initialize a stream wrapper that redirects writes to a callback while preserving the original."""
            self.original = original
            self.callback = callback

        def write(self, text):
            """Write text to the original stream and forward it to the capture callback."""
            self.original.write(text)
            if self.callback:
                self.callback(text, None)

        def progress(self, bar_id, frame):
            """Write a TQDM frame to the original stream and report it as bar state instead of a log line."""
            self.original.write(frame)
            if self.callback:
                self.callback(frame, bar_id)

        def flush(self):
            """Flush the wrapped stream to propagate buffered output promptly during console capture."""
            self.original.flush()

        def isatty(self):
            """Delegate isatty check to the original stream."""
            return self.original.isatty()

    class _LogHandler(logging.Handler):
        """Lightweight logging handler."""

        __slots__ = ("callback",)

        def __init__(self, callback):
            """Initialize a lightweight logging.Handler that forwards log records to the provided callback."""
            super().__init__()
            self.callback = callback

        def emit(self, record):
            """Format and forward LogRecord messages to the capture callback for unified log streaming."""
            self.callback(self.format(record) + "\n", None)


class _DriveInfo:
    """Resolve mounted storage paths backed by local drives.

    This helper keeps platform-specific drive discovery isolated from SystemLogger metric collection. It uses fast
    psutil mount discovery first and falls back to native OS commands only when multiple visible mounts need
    disambiguation.

    Examples:
        >>> logger = SystemLogger()
        >>> logger.mounts
        ['/']
    """

    @staticmethod
    def mounts(psutil, all_drives=False):
        """Get mounted paths to monitor."""
        partitions = [p for p in psutil.disk_partitions(all=False) if p.mountpoint]
        if not all_drives:
            return [_DriveInfo._current_mount(partitions)]

        mounts = [
            p.mountpoint for p in partitions if Path(p.mountpoint).is_dir() and "dontbrowse" not in p.opts.split(",")
        ]
        if len(mounts) <= 1:
            return _DriveInfo._sort(mounts) or [_DriveInfo._current_mount(partitions)]

        for getter in (
            _DriveInfo._macos_mounts if MACOS else None,
            _DriveInfo._linux_mounts if LINUX else None,
            _DriveInfo._windows_mounts if WINDOWS else None,
        ):
            if getter:
                try:
                    if platform_mounts := getter(partitions):
                        return _DriveInfo._sort(platform_mounts)
                except (json.JSONDecodeError, OSError, plistlib.InvalidFileException, subprocess.SubprocessError):
                    pass
        return _DriveInfo._sort(mounts)

    @staticmethod
    def _sort(mounts):
        """Sort mounted paths with root first, excluding boot/firmware partitions like /boot and /boot/efi."""
        mounts = {m for m in mounts if not (m + "/").startswith(("/boot/", "/efi/"))}
        return sorted(mounts, key=lambda mount: (mount != "/", mount))

    @staticmethod
    def _current_mount(partitions):
        """Get the mounted filesystem backing the current working directory."""
        try:
            cwd = Path.cwd().resolve()
        except OSError:
            return "C:\\" if WINDOWS else "/"
        matches = []
        for partition in partitions:
            try:
                mount = Path(partition.mountpoint).resolve()
            except OSError:
                continue
            if cwd == mount or cwd.is_relative_to(mount):
                matches.append(partition.mountpoint)
        return max(matches, key=len, default=Path.cwd().anchor or "/")

    @staticmethod
    def _macos_mounts(partitions):
        """Get user-visible macOS mounts backed by physical disks."""
        disk_info = plistlib.loads(subprocess.check_output(["diskutil", "list", "-plist", "physical"], timeout=5))
        physical_devices = set(disk_info.get("WholeDisks", []))
        for disk in disk_info.get("AllDisksAndPartitions", []):
            physical_devices.add(disk.get("DeviceIdentifier", ""))
            physical_devices.update(p.get("DeviceIdentifier", "") for p in disk.get("Partitions", []))

        mounts, volume_groups = [], set()
        for partition in partitions:
            if partition.mountpoint != "/" and "dontbrowse" in partition.opts.split(","):
                continue
            info = plistlib.loads(
                subprocess.check_output(
                    ["diskutil", "info", "-plist", partition.mountpoint],
                    stderr=subprocess.DEVNULL,
                    timeout=5,
                )
            )
            devices = {info.get("DeviceIdentifier", "")}
            devices.update(s.get("APFSPhysicalStore", "") for s in info.get("APFSPhysicalStores", []))
            if not devices & physical_devices:
                continue
            group = info.get("APFSVolumeGroupID") or info.get("APFSContainerReference") or partition.mountpoint
            if group in volume_groups:
                continue
            volume_groups.add(group)
            mounts.append(partition.mountpoint)
        return mounts

    @staticmethod
    def _linux_mounts(_partitions):
        """Get Linux mounts backed by physical block devices."""
        block_info = json.loads(
            subprocess.check_output(
                ["lsblk", "--json", "--output", "NAME,TYPE,MOUNTPOINT,MOUNTPOINTS"], text=True, timeout=5
            )
        )
        mounts = []

        def visit(block, physical=False):
            physical = physical or block.get("type") == "disk"
            if physical:
                values = block.get("mountpoints") or [block.get("mountpoint")]
                if isinstance(values, str):
                    values = [values]
                mounts.extend(m for m in values if isinstance(m, str) and m.startswith("/") and Path(m).is_dir())
            for child in block.get("children", []):
                visit(child, physical)

        for block in block_info.get("blockdevices", []):
            visit(block)
        return mounts

    @staticmethod
    def _windows_mounts(_partitions):
        """Get Windows fixed local drive mounts."""
        output = subprocess.check_output(
            [
                "powershell",
                "-NoProfile",
                "-Command",
                "Get-CimInstance Win32_LogicalDisk -Filter 'DriveType=3' | Select-Object -ExpandProperty DeviceID",
            ],
            text=True,
            timeout=5,
        )
        return [f"{drive}\\" for drive in (line.strip() for line in output.splitlines()) if drive]


class SystemLogger:
    """Log dynamic system metrics for training monitoring.

    Captures real-time system metrics including CPU, RAM, disk I/O, network I/O, and NVIDIA GPU statistics for training
    performance monitoring and analysis.

    Attributes:
        pynvml (module | None): NVIDIA pynvml module if successfully imported, None otherwise.
        nvidia_initialized (bool): Whether NVIDIA GPU monitoring is available and initialized.
        nvidia_versions (dict): NVIDIA 'driver_version' and 'cuda_version' strings read from NVML, if available.
        net_start (namedtuple): Initial network I/O counters for calculating cumulative usage.
        disk_start (namedtuple | None): Initial disk I/O counters for calculating cumulative usage.
        mounts (list[str]): Mounted drive paths monitored for disk usage.

    Examples:
        Basic usage (single drive):
        >>> logger = SystemLogger()
        >>> metrics = logger.get_metrics()
        >>> print(f"CPU: {metrics['cpu']}%, RAM: {metrics['ram']}%")
        >>> for disk in metrics["disk"]:
        ...     print(f"{disk['mount']}: {disk['used_gb']}/{disk['total_gb']} GB")

        Monitor all drives:
        >>> logger = SystemLogger(all_drives=True)
        >>> metrics = logger.get_metrics()
        >>> for disk in metrics["disk"]:
        ...     print(f"{disk['mount']}: {disk['used_gb']}/{disk['total_gb']} GB")

        Training loop integration:
        >>> system_logger = SystemLogger()
        >>> for epoch in range(epochs):
        ...     # Training code here
        ...     metrics = system_logger.get_metrics()
        ...     # Log to database/file
    """

    def __init__(self, all_drives=False):
        """Initialize the system logger.

        Args:
            all_drives (bool): If True, monitor all mounted drives. If False, monitor the current drive.
        """
        import psutil  # scoped as slow import

        self.pynvml = None
        self.nvidia_initialized = self._init_nvidia()
        self.nvidia_versions = {}
        if self.nvidia_initialized:
            try:
                self.nvidia_versions["driver_version"] = self.pynvml.nvmlSystemGetDriverVersion()
                cuda = self.pynvml.nvmlSystemGetCudaDriverVersion_v2()
                self.nvidia_versions["cuda_version"] = f"{cuda // 1000}.{cuda % 1000 // 10}"
            except self.pynvml.NVMLError:
                pass
        self.net_start = psutil.net_io_counters()
        self.disk_start = psutil.disk_io_counters()
        self.mounts = _DriveInfo.mounts(psutil, all_drives)

        # For rate calculation
        self._prev_net = self.net_start
        self._prev_disk = self.disk_start
        self._prev_time = time.time()

    def _init_nvidia(self):
        """Initialize NVIDIA GPU monitoring with pynvml."""
        if MACOS:
            return False

        try:
            import pynvml  # scoped as slow import

            self.pynvml = pynvml
            pynvml.nvmlInit()
            return True
        except Exception as e:
            import torch

            if torch.cuda.is_available():
                LOGGER.warning(f"SystemLogger NVML init failed: {e}")
            return False

    def get_metrics(self, rates=False):
        """Get current system metrics including CPU, RAM, disk, network, and GPU usage.

        Collects comprehensive system metrics including CPU usage, RAM usage, disk usage, disk I/O statistics, network
        I/O statistics, and GPU metrics (if available).

        On NVIDIA systems, also reports `driver_version` and `cuda_version` cached from NVML at initialization. CUDA is the
        driver-supported version shown by nvidia-smi, not the installed toolkit or PyTorch build version.

        Example output (rates=False, default):
        ```python
        {
            "cpu": 45.2,
            "ram": 78.9,
            "disk": [{"mount": "/", "used_gb": 256.8, "total_gb": 512.0}],
            "disk_io": {"read_mb": 156.7, "write_mb": 89.3},
            "network": {"recv_mb": 157.2, "sent_mb": 89.1},
            "gpus": {
                "0": {"usage": 95.6, "memory": 85.4, "temp": 72, "power": 285},
                "1": {"usage": 94.1, "memory": 82.7, "temp": 70, "power": 278},
            },
        }
        ```

        Example output (rates=True):
        ```python
        {
            "cpu": 45.2,
            "ram": 78.9,
            "disk": [{"mount": "/", "used_gb": 256.8, "total_gb": 512.0}],
            "disk_io": {"read_mbs": 12.5, "write_mbs": 8.3},
            "network": {"recv_mbs": 5.2, "sent_mbs": 1.1},
            "gpus": {
                "0": {"usage": 95.6, "memory": 85.4, "temp": 72, "power": 285},
            },
        }
        ```

        Args:
            rates (bool): If True, return disk/network as MB/s rates instead of cumulative MB.

        Returns:
            (dict): Metrics dictionary with cpu, ram, disk, disk_io, network, and gpus keys, plus driver_version and
                cuda_version on NVIDIA systems.

        Examples:
            >>> logger = SystemLogger()
            >>> logger.get_metrics()["cpu"]  # CPU percentage
            >>> logger.get_metrics(rates=True)["network"]["recv_mbs"]  # MB/s download rate
        """
        import psutil  # scoped as slow import

        net = psutil.net_io_counters()
        disk_io = psutil.disk_io_counters()
        memory = psutil.virtual_memory()
        now = time.time()

        # Calculate elapsed time since last call
        elapsed = max(0.1, now - self._prev_time)  # Avoid division by zero

        if rates:
            disk_io_metrics = {
                "read_mbs": round(max(0, (disk_io.read_bytes - self._prev_disk.read_bytes) / 1e6 / elapsed), 3),
                "write_mbs": round(max(0, (disk_io.write_bytes - self._prev_disk.write_bytes) / 1e6 / elapsed), 3),
            }
        else:
            disk_io_metrics = {
                "read_mb": round((disk_io.read_bytes - self.disk_start.read_bytes) / 1e6, 3),
                "write_mb": round((disk_io.write_bytes - self.disk_start.write_bytes) / 1e6, 3),
            }

        disks = []
        for mounts in (self.mounts, ["C:\\" if WINDOWS else "/"]):
            for mount in mounts:
                try:
                    usage = shutil.disk_usage(mount)
                    disks.append(
                        {
                            "mount": mount,
                            "used_gb": round(usage.used / 1e9, 3),
                            "total_gb": round(usage.total / 1e9, 3),
                        }
                    )
                except (PermissionError, OSError):
                    continue  # Skip inaccessible drives
            if disks:
                break

        metrics = {
            **self.nvidia_versions,
            "cpu": round(psutil.cpu_percent(), 3),
            "ram": round(memory.percent, 3),
            "disk": disks,
            "disk_io": disk_io_metrics,
            "gpus": {},
        }

        if rates:
            metrics["network"] = {
                "recv_mbs": round(max(0, (net.bytes_recv - self._prev_net.bytes_recv) / 1e6 / elapsed), 3),
                "sent_mbs": round(max(0, (net.bytes_sent - self._prev_net.bytes_sent) / 1e6 / elapsed), 3),
            }
        else:
            metrics["network"] = {
                "recv_mb": round((net.bytes_recv - self.net_start.bytes_recv) / 1e6, 3),
                "sent_mb": round((net.bytes_sent - self.net_start.bytes_sent) / 1e6, 3),
            }

        # Always update previous values for accurate rate calculation on next call
        self._prev_net = net
        self._prev_disk = disk_io
        self._prev_time = now

        # Add GPU metrics (NVIDIA only)
        if self.nvidia_initialized:
            metrics["gpus"].update(self._get_nvidia_metrics())

        return metrics

    def _get_nvidia_metrics(self):
        """Get NVIDIA GPU metrics including utilization, memory, temperature, and power."""
        gpus = {}
        if not self.nvidia_initialized or not self.pynvml:
            return gpus
        try:
            device_count = self.pynvml.nvmlDeviceGetCount()
            for i in range(device_count):
                handle = self.pynvml.nvmlDeviceGetHandleByIndex(i)
                util = self.pynvml.nvmlDeviceGetUtilizationRates(handle)
                memory = self.pynvml.nvmlDeviceGetMemoryInfo(handle)
                temp = self.pynvml.nvmlDeviceGetTemperature(handle, self.pynvml.NVML_TEMPERATURE_GPU)
                power = self.pynvml.nvmlDeviceGetPowerUsage(handle) // 1000

                gpus[str(i)] = {
                    "usage": round(util.gpu, 3),
                    "memory": round((memory.used / memory.total) * 100, 3),
                    "temp": temp,
                    "power": power,
                }
        except Exception:
            pass
        return gpus


if __name__ == "__main__":
    print("SystemLogger Real-time Metrics Monitor")
    print("Press Ctrl+C to stop\n")

    logger = SystemLogger(all_drives=True)

    try:
        while True:
            metrics = logger.get_metrics()

            # Clear screen (works on most terminals)
            print("\033[H\033[J", end="", flush=True)

            # Display system metrics
            print(f"CPU: {metrics['cpu']:5.1f}%")
            print(f"RAM: {metrics['ram']:5.1f}%")
            print(f"Net Recv: {metrics['network']['recv_mb']:9.1f} MB")
            print(f"Net Sent: {metrics['network']['sent_mb']:9.1f} MB")

            # Display disk metrics
            print("\nDisk Metrics:")
            for disk in metrics["disk"]:
                print(f"  {disk['mount']}: {disk['used_gb']:.1f}/{disk['total_gb']:.1f} GB")

            # Display GPU metrics if available
            if metrics["gpus"]:
                print("\nGPU Metrics:")
                for gpu_id, gpu_data in metrics["gpus"].items():
                    print(
                        f"  GPU {gpu_id}: {gpu_data['usage']:3}% | "
                        f"Mem: {gpu_data['memory']:5.1f}% | "
                        f"Temp: {gpu_data['temp']:2}°C | "
                        f"Power: {gpu_data['power']:3}W"
                    )
            else:
                print("\nGPU: No NVIDIA GPUs detected")

            time.sleep(1)

    except KeyboardInterrupt:
        print("\n\nStopped monitoring.")
