"""Analytic baselines on RadioMapSeer.

B0  free space: PL = a + b*log10(d) (fitted log-distance; slope ~20 is free space)
B1  multi-wall (COST-231-style, straight line): B0 + c*metres_inside_buildings + e*wall_crossings + g*LOS
B2  street canyon (b2.py): a + b*log10(geodesic d) + c*turns + e*LOS + g*metres_inside + h*blocked
All are fitted by least squares on training pixels above the analytic floor and then clipped
to the dataset range. B1 is the "physics channel" fed to the hybrid network.
"""
from __future__ import annotations

import json
from pathlib import Path

import numpy as np

from .b2 import UNREACHED, turns_from_cache
from .data import DB_RANGE, N, P_TRNC

FEATS = {"B0": ["one", "logd"], "B1": ["one", "logd", "inside", "walls", "los"],
         "B2": ["one", "loggeod", "turns", "los", "inside", "blocked"]}
EXTRA = {"B0": [], "B1": [], "B2": ["geod", "turns"]}  # cache arrays needed beyond inside/walls


def design(inside, walls, r, c, names, geod=None, turns=None):
    rr, cc = np.mgrid[:inside.shape[0], :inside.shape[1]]
    d = np.maximum(np.hypot(rr - r, cc - c), 1.0)
    cols = {"one": np.ones(inside.shape, np.float32), "logd": np.log10(d).astype(np.float32),
            "inside": inside.astype(np.float32), "walls": walls.astype(np.float32),
            "los": (inside == 0).astype(np.float32)}
    if geod is not None:
        blocked = geod == UNREACHED
        cols["loggeod"] = np.where(blocked, 0.0, np.log10(np.maximum(geod, 1.0))).astype(np.float32)
        cols["blocked"] = blocked.astype(np.float32)
        cols["turns"] = turns_from_cache(turns.astype(np.float32))
    return np.stack([cols[n] for n in names], -1)


def fit(cache: Path, name: str = "B1", target: str = "dpm", n_samples: int = 400, px: int = 2000, seed: int = 0,
        maps_limit: int | None = None):
    """Least squares on pixels above the floor of n_samples random training maps (only the first
    maps_limit training maps if given, so a data-limited run does not see the other maps)."""
    rng = np.random.default_rng(seed)
    meta = json.loads((cache / "train_meta.json").read_text())
    pool = np.arange(len(meta))
    if maps_limit:
        keep = set(json.loads((cache / "train_maps.json").read_text())[:maps_limit])
        pool = np.array([i for i in pool if meta[i]["map"] in keep])
    ins = np.load(cache / "train_inside.npy", mmap_mode="r")
    wal = np.load(cache / "train_walls.npy", mmap_mode="r")
    tgt = np.load(cache / f"train_{target}.npy", mmap_mode="r")
    ex = {k: np.load(cache / f"train_{k}.npy", mmap_mode="r") for k in EXTRA[name]}
    X, Y = [], []
    for i in rng.choice(pool, min(n_samples, len(pool)), replace=False):
        m = meta[i]
        D = design(ins[i], wal[i], m["r"], m["c"], FEATS[name], **{k: a[i] for k, a in ex.items()}).reshape(-1, len(FEATS[name]))
        y = tgt[i].reshape(-1).astype(np.float32)
        ok = np.flatnonzero(y > 0)  # above the analytic floor (truncation would bias the fit)
        sel = rng.choice(ok, min(px, len(ok)), replace=False)
        X.append(D[sel]); Y.append(P_TRNC + DB_RANGE * y[sel] / 255.0)
    X, Y = np.concatenate(X), np.concatenate(Y)
    coef, *_ = np.linalg.lstsq(X, Y, rcond=None)
    return coef


def predict_gray(coef, inside, walls, r, c, name="B1", geod=None, turns=None, db_out=False):
    """Physics map in gray (dataset encoding, clipped); db_out=True: unclipped path gain in dB instead."""
    db = design(inside, walls, r, c, FEATS[name], geod, turns) @ coef
    if db_out:
        return db.astype(np.float32)
    return np.clip((db - P_TRNC) / DB_RANGE, 0.0, 1.0).astype(np.float32)
