from __future__ import annotations

import logging
import threading
from datetime import datetime

_tl = threading.local()

_MAX_OUTPUT_BYTES = 100 * 1024  # 100 KB cap

_active_count = 0
_active_count_lock = threading.Lock()


class CeleryTaskLogHandler(logging.Handler):
    """
    A logging.Handler that buffers log records per Celery task execution.

    Activated via attach_handler(); records are buffered on thread-local storage
    and retrieved via flush_and_detach() on task completion.
    """

    def emit(self, record: logging.LogRecord) -> None:
        if not getattr(_tl, "active_task_id", None):
            return
        buf: list[str] = getattr(_tl, "log_buffer", None)
        if buf is None:
            return
        # Format: [LEVEL HH:MM:SS name] message
        ts = datetime.utcnow().strftime("%H:%M:%S")
        line = f"[{record.levelname} {ts} {record.name}] {self.format(record)}"
        buf.append(line)

    def format(self, record: logging.LogRecord) -> str:  # type: ignore[override]
        return record.getMessage()


_GLOBAL_HANDLER = CeleryTaskLogHandler()
_GLOBAL_HANDLER.setLevel(logging.DEBUG)


def attach_handler(task_id: str, log_id: int) -> None:
    """Mark the current thread as belonging to task_id and start buffering."""
    global _active_count
    _tl.active_task_id = task_id
    _tl.active_log_id = log_id
    _tl.log_buffer = []
    with _active_count_lock:
        _active_count += 1
        if _GLOBAL_HANDLER not in logging.getLogger().handlers:
            logging.getLogger().addHandler(_GLOBAL_HANDLER)


def flush_and_detach(task_id: str) -> str:
    """
    Return the buffered log output for task_id and clear thread-local state.
    Removes the handler from the root logger only when no other tasks are running.
    """
    global _active_count
    if getattr(_tl, "active_task_id", None) != task_id:
        return ""

    lines: list[str] = getattr(_tl, "log_buffer", [])
    output = "\n".join(lines)

    # Truncate at 100 KB
    if len(output.encode("utf-8")) > _MAX_OUTPUT_BYTES:
        output = output.encode("utf-8")[:_MAX_OUTPUT_BYTES].decode("utf-8", errors="ignore")
        output += "\n... [output truncated at 100 KB]"

    _tl.active_task_id = None
    _tl.active_log_id = None
    _tl.log_buffer = []

    with _active_count_lock:
        _active_count = max(0, _active_count - 1)
        if _active_count == 0:
            try:
                logging.getLogger().removeHandler(_GLOBAL_HANDLER)
            except Exception:
                pass

    return output
