import logging
import os
from datetime import datetime

LOG_DIR = os.getenv("LOG_DIR", "logs")
_FMT = "%(asctime)s [%(levelname)s] %(name)s: %(message)s"


class MonthlyFileHandler(logging.Handler):
    """Writes to logs/YYYY-MM[-error].log, rotates when month changes."""

    def __init__(self, log_dir: str = LOG_DIR, errors_only: bool = False):
        super().__init__()
        self.log_dir = log_dir
        self.errors_only = errors_only
        self._month: str | None = None
        self._stream = None

    def _suffix(self) -> str:
        return "-error" if self.errors_only else ""

    def _current_stream(self):
        month = datetime.now().strftime("%Y-%m")
        if month != self._month:
            if self._stream:
                self._stream.close()
            os.makedirs(self.log_dir, exist_ok=True)
            self._stream = open(
                os.path.join(self.log_dir, f"{month}{self._suffix()}.log"),
                "a", encoding="utf-8"
            )
            self._month = month
        return self._stream

    def emit(self, record):
        try:
            msg = self.format(record)
            s = self._current_stream()
            s.write(msg + "\n")
            s.flush()
        except Exception:
            self.handleError(record)

    def close(self):
        if self._stream:
            self._stream.close()
        super().close()


def setup_logging(level: int | None = None) -> None:
    import os
    if level is None:
        level = logging.DEBUG if os.getenv("DEBUG", "").lower() == "true" else logging.INFO
    fmt = logging.Formatter(_FMT)

    console = logging.StreamHandler()
    console.setFormatter(fmt)

    # logs/YYYY-MM.log — INFO and above, excludes ERROR+
    general_h = MonthlyFileHandler(errors_only=False)
    general_h.setFormatter(fmt)
    general_h.addFilter(lambda r: r.levelno < logging.ERROR)

    # logs/YYYY-MM-error.log — ERROR and CRITICAL only
    error_h = MonthlyFileHandler(errors_only=True)
    error_h.setFormatter(fmt)
    error_h.setLevel(logging.ERROR)

    root = logging.getLogger()
    root.setLevel(level)
    if not root.handlers:
        root.addHandler(console)
        root.addHandler(general_h)
        root.addHandler(error_h)

    # suppress noisy third-party debug output
    for noisy in ("httpcore", "httpx", "google_genai", "openai", "anthropic", "urllib3"):
        logging.getLogger(noisy).setLevel(logging.WARNING)
