"""Per-model preset generation (--models-preset INI) — the router-side carrier for context-policy
launch decisions.
"""

from __future__ import annotations

import json
import logging
from dataclasses import dataclass, replace
from pathlib import Path

from hermes_cli.local_runtime.context_policy import (
    RUNTIME_OVERHEAD_BYTES, fit_to_free_memory, launch_args, plan_launch, ub_logits_bytes)
from hermes_cli.local_runtime.estimator import (
    HardwareBudget, PhysicsRefusal, ctx_bytes, footprint_bytes, profile_from_gguf)
from hermes_cli.local_runtime.gguf import model_id_from_stem, read_gguf_header

logger = logging.getLogger(__name__)

# args list -> INI keys. Flags the policy owns; everything else stays out of the preset.
_FLAG_TO_KEY = {
    "-c": "ctx-size", "-b": "batch-size", "-ub": "ubatch-size",
    "-ctk": "cache-type-k", "-ctv": "cache-type-v", "-fa": "flash-attn",
    "-ot": "override-tensor", "--spec-type": "spec-type", "--spec-draft-n-max": "spec-draft-n-max",
}


@dataclass
class PresetEntry:
    model_id: str
    window: int
    spilled: bool
    refusal: str | None = None
    keys: dict[str, str] | None = None


def _args_to_keys(args: list[str]) -> dict[str, str]:
    keys: dict[str, str] = {}
    i = 0
    while i < len(args):
        key = _FLAG_TO_KEY.get(args[i])
        if key is None:
            i += 1
            continue
        keys[key] = args[i + 1]
        i += 2
    return keys


def _asset_path(asset) -> "Path | None":
    """On-disk path of a catalog companion asset, or None when it isn't downloaded."""
    from hermes_cli.local_runtime.bootstrap import assets_dir

    if asset is None:
        return None
    path = assets_dir() / asset.local_name
    return path if path.exists() else None


def _draft_fits(path: Path, profile, budget: HardwareBudget, window: int, overhead: int) -> bool:
    """Optional draft never shrinks the advertised window or displaces its GPU buffers.

    The catalog does not know the draft's context layout. Admit it only after reading the file;
    draft KV defaults to f16, independently of the target's q8 cache.
    """
    try:
        draft = profile_from_gguf(read_gguf_header(path))
        draft_need = footprint_bytes(
            draft, window, flash_attention=False,
            overhead_bytes=RUNTIME_OVERHEAD_BYTES + ub_logits_bytes(draft.n_vocab, mtp_capable=False))
    except (ValueError, OSError) as exc:
        logger.warning("draft omitted %s: %s", path.name, exc)
        return False
    return (footprint_bytes(profile, window, overhead_bytes=overhead) + draft_need
            <= budget.usable_vram_bytes + budget.ram_available_bytes
            and ctx_bytes(profile, window) + overhead + draft_need <= budget.usable_vram_bytes)


def preset_for_model(gguf: Path, budget: HardwareBudget,
                     mtp_capable: set[str], *, requested_window: int | None = None,
                     live: HardwareBudget | None = None) -> PresetEntry | None:
    """The launch decision for one staged model, or None when its header is unreadable.

    ``budget`` is the card's capacity; ``live`` (from ``hardware.launch_budget``) narrows the
    window to what fits beside other programs now.
    """
    from hermes_cli.local_runtime.catalog import entry_for_model
    from hermes_cli.local_runtime.growth import load_window_overrides

    model_id = model_id_from_stem(gguf.stem)
    try:
        header = read_gguf_header(gguf)
        profile = profile_from_gguf(header)
    except (ValueError, OSError) as exc:
        logger.warning("preset skip %s: %s", gguf.name, exc)
        return None
    entry = entry_for_model(model_id)
    is_mtp = entry.mtp if entry is not None else model_id in mtp_capable

    mmproj_path = _asset_path(entry.mmproj) if entry is not None else None
    fixed_overhead = RUNTIME_OVERHEAD_BYTES + (
        entry.mmproj.size_bytes if entry is not None and mmproj_path is not None else 0)
    plan = plan_launch(profile, budget, mtp_capable=is_mtp, fixed_overhead=fixed_overhead,
                       requested_window=(load_window_overrides().get(model_id)
                                         if requested_window is None else requested_window))
    if live is not None:
        plan = fit_to_free_memory(plan, profile, live, mtp_capable=is_mtp,
                                  fixed_overhead=fixed_overhead)
    decision = plan.decision
    if isinstance(decision, PhysicsRefusal):
        return PresetEntry(model_id=model_id, window=0, spilled=False, refusal=decision.message)

    # Router discovery is preset-only: refused files must never autoload with stock fit.
    keys = _args_to_keys(launch_args(
        profile, decision, mtp_capable=is_mtp, uma=budget.uma, mtp_prefill=plan.mtp_prefill,
        mtp_draft_depth=entry.mtp_draft_depth if entry is not None else 3))
    keys["model"] = str(gguf)
    if entry is not None and is_mtp:
        # Integrated-MTP targets sample on the backend, and so does the draft (pairing validated
        # against the vendor's published llama.cpp recipes).
        keys["backend-sampling"] = "on"
        keys["spec-draft-backend-sampling"] = "on"

    # Sampling deference ladder, under the policy keys (policy wins on clash): the GGUF's own
    # general.sampling.* metadata is the publisher's recommendation and covers models the catalog
    # has never heard of; catalog sampling applies only where the file is silent; a model
    # carrying neither runs llama.cpp defaults.
    for k, v in header.sampling_defaults.items():
        keys.setdefault(k, v)
    if entry is not None:
        for k, v in (entry.sampling or {}).items():
            keys.setdefault(k, v)
        if mmproj_path is not None:
            keys["mmproj"] = str(mmproj_path)
        draft_path = _asset_path(entry.draft) if decision.spilled else None
        if draft_path is not None and _draft_fits(draft_path, profile, budget, decision.window, plan.overhead_bytes):
            keys["model-draft"] = str(draft_path)
            keys["spec-type"] = "draft-dspark"
            # Unsloth's measured cliff: acceptance 83% at 2-3 drafts, collapses at 4.
            keys["spec-draft-n-max"] = "3"
    return PresetEntry(model_id=model_id, window=decision.window,
                       spilled=decision.spilled, keys=keys)


def resident_footprint(gguf: Path, budget: HardwareBudget, window: int) -> int | None:
    """Estimated bytes one staged model holds while loaded at ``window``, or None when unreadable."""
    from hermes_cli.local_runtime.catalog import entry_for_model

    model_id = model_id_from_stem(gguf.stem)
    try:
        profile = profile_from_gguf(read_gguf_header(gguf))
    except (ValueError, OSError) as exc:
        logger.debug("footprint skip %s: %s", gguf.name, exc)
        return None
    entry = entry_for_model(model_id)
    is_mtp = entry.mtp if entry is not None else False
    mmproj = entry.mmproj.size_bytes if entry is not None and _asset_path(entry.mmproj) else 0
    plan = plan_launch(profile, budget, mtp_capable=is_mtp,
                       fixed_overhead=RUNTIME_OVERHEAD_BYTES + mmproj, requested_window=window)
    if is_mtp and profile.kv_scale == 1.0:
        profile = replace(profile, kv_scale=1.2)
    return footprint_bytes(profile, window, overhead_bytes=plan.overhead_bytes)


def _launch_footprint(gguf: Path, budget: HardwareBudget) -> int | None:
    """Estimated resident bytes for one staged model at the window this policy grants it, or None
    when it cannot be priced: an unreadable header, or a model the physics check refuses outright
    (it never loads, so it must not shrink the residency cap)."""
    from hermes_cli.local_runtime.catalog import entry_for_model
    from hermes_cli.local_runtime.growth import load_window_overrides

    model_id = model_id_from_stem(gguf.stem)
    try:
        profile = profile_from_gguf(read_gguf_header(gguf))
    except (ValueError, OSError) as exc:
        logger.debug("footprint skip %s: %s", gguf.name, exc)
        return None
    entry = entry_for_model(model_id)
    is_mtp = entry.mtp if entry is not None else False
    mmproj = entry.mmproj.size_bytes if entry is not None and _asset_path(entry.mmproj) else 0
    plan = plan_launch(profile, budget, mtp_capable=is_mtp,
                       fixed_overhead=RUNTIME_OVERHEAD_BYTES + mmproj,
                       requested_window=load_window_overrides().get(model_id))
    if isinstance(plan.decision, PhysicsRefusal):
        return None
    # Priced whole even when the plan spills: a spilled model still holds part of its weights on
    # the device, and over-counting errs toward the side that cannot thrash.
    return footprint_bytes(profile, plan.decision.window, overhead_bytes=plan.overhead_bytes)


def admitted_residency_count(models_dir: Path, budget: HardwareBudget, configured: int) -> int:
    """How many models the card may hold resident at once: priced against the budget, not a count.

    Residency used to be bounded by a count alone, so a second model was admitted against an
    already-full card. On Windows/WDDM that over-commit is not refused — the allocation is paged
    to host memory, and that child decodes at a third of its speed for the rest of its life: no
    error, no UI hint, and ejecting the incumbent afterwards does not repair it (only a clean
    reload does). Capping the count instead has llama.cpp evict its LRU *before* the incoming
    child allocates, which is the only placement that fits.

    The cap rises above one only while the LARGEST staged model still fits TWICE — any pair of
    staged models then fits by construction. ``configured`` stays a ceiling (a user's smaller
    number is honoured), and an unpriceable input (no usable device memory, no readable model)
    keeps today's behaviour.
    """
    from hermes_cli.local_runtime.bootstrap import staged_in

    if configured <= 1 or budget.usable_vram_bytes <= 0:
        return configured
    largest = 0
    for gguf in staged_in(models_dir):
        need = _launch_footprint(gguf, budget)
        if need:
            largest = max(largest, need)
    if largest <= 0:
        return configured
    return max(1, min(configured, budget.usable_vram_bytes // largest))


def plan_presets(models_dir: Path, budget: HardwareBudget, mtp_capable: set[str] | None = None,
                 *, live: HardwareBudget | None = None) -> list[PresetEntry]:
    """The launch decision for every staged model; unreadable headers are skipped."""
    from hermes_cli.local_runtime.bootstrap import staged_in

    entries = (preset_for_model(gguf, budget, mtp_capable or set(), live=live)
               for gguf in staged_in(models_dir))
    return [entry for entry in entries if entry is not None]


def generate_presets(models_dir: Path, budget: HardwareBudget, preset_path: Path,
                     mtp_capable: set[str] | None = None, *,
                     live: HardwareBudget | None = None) -> list[PresetEntry]:
    """Plan every staged model and write one INI. Refused models get no section (the picker
    surfaces the refusal from the returned entries)."""
    entries = plan_presets(models_dir, budget, mtp_capable, live=live)
    write_presets(entries, preset_path)
    return entries


def write_presets(entries: list[PresetEntry], preset_path: Path) -> None:
    from utils import atomic_write_text

    atomic_write_text(preset_path, render_presets(entries), tmp_prefix=f".{preset_path.name}_", mode=0o600)
    logger.info("wrote %d preset sections to %s", sum(e.keys is not None for e in entries), preset_path)


def render_presets(entries: list[PresetEntry]) -> str:
    sections: list[str] = []
    for entry in entries:
        # INI comments preserve non-flag facts atomically with the launch policy.
        sections.append("# hermes-decision: " + json.dumps({
            "model_id": entry.model_id, "window": entry.window,
            "spilled": entry.spilled, "refusal": entry.refusal}) + "\n")
        if entry.keys is not None:
            body = "\n".join(f"{k} = {v}" for k, v in entry.keys.items())
            sections.append(f"[{entry.model_id}]\n{body}\n")
    return "\n".join(sections)


def read_preset_decisions(preset_path: Path | None = None) -> dict[str, PresetEntry]:
    """The launch decisions the running server was actually given, read back from the preset INI
    (the INI is the record — it's what spawned the children). Missing/unparseable -> {}."""
    import configparser

    if preset_path is None:
        from hermes_cli.local_runtime.binaries import runtimes_root

        preset_path = runtimes_root() / "presets.ini"
    out: dict[str, PresetEntry] = {}
    try:
        parser = configparser.ConfigParser(interpolation=None)
        text = preset_path.read_text(encoding="utf-8-sig")
        parser.read_string(text)
        recorded = {}
        for line in text.splitlines():
            if line.startswith("# hermes-decision: "):
                fact = json.loads(line.removeprefix("# hermes-decision: "))
                recorded[fact["model_id"]] = fact
                if fact.get("refusal"):
                    out[fact["model_id"]] = PresetEntry(**fact)
        for section in parser.sections():
            out[section] = PresetEntry(
                model_id=section, window=parser.getint(section, "ctx-size", fallback=0),
                spilled=recorded.get(section, {}).get("spilled", parser.has_option(section, "override-tensor")),
                keys=dict(parser[section]))
    except Exception as exc:  # noqa: BLE001
        logger.debug("preset read-back failed: %s", exc)
    return out
