"""SkyDiscover GEPA Native population: full archive, elite pool, and rejection history."""
from __future__ import annotations

import collections
from typing import Any

from ...checkpoint import genome_from_dict, genome_to_dict
from ...components.population import Population
from ...records import Genome


def score_metrics(metrics: dict[str, Any] | None) -> float:
    """Mirror SkyDiscover ``get_score`` exactly for native algorithm decisions."""
    if not metrics:
        return 0.0
    if "combined_score" in metrics:
        try:
            return float(metrics["combined_score"])
        except (TypeError, ValueError):
            pass
    numeric = [
        value
        for value in metrics.values()
        if isinstance(value, (int, float)) and not isinstance(value, bool)
    ]
    return float(sum(numeric) / len(numeric)) if numeric else 0.0


def genome_score(genome: Genome | None) -> float:
    return score_metrics(genome.scores if genome is not None else None)


class GepaNativePopulation(Population):
    """Accepted full archive plus a fixed-size sampling pool.

    Normal children are admitted only when they strictly improve over their stored parent (unless the
    gate is disabled).  Rejected normal mutations go only to ``rejection_history``.  Merge children
    use the separate upstream criterion ``score >= max(parent_a, parent_b)`` and merge rejections are
    deliberately not reflective-history entries.
    """

    def __init__(
        self,
        population_size: int = 40,
        max_rejection_history: int = 20,
        acceptance_gating: bool = True,
    ):
        self.population_size = int(population_size)
        self.acceptance_gating = bool(acceptance_gating)
        self.max_rejection_history = int(max_rejection_history)
        self._programs: dict[str, Genome] = {}
        self.elite_pool: list[str] = []
        self.rejection_history: collections.deque[Genome] = collections.deque(
            maxlen=self.max_rejection_history
        )
        self.metric_best: dict[str, tuple[str, float]] = {}
        self.program_at_metric_front: dict[str, set[str]] = {}
        self.initial_program_id: str | None = None
        self.best_program_id: str | None = None

    @property
    def programs(self) -> dict[str, Genome]:
        """Read-only-by-convention view matching the upstream database name."""
        return self._programs

    def get(self, program_id: str | None) -> Genome | None:
        return self._programs.get(program_id) if program_id else None

    def _merge_is_acceptable(self, genome: Genome) -> bool:
        score_a = genome.metadata.get("merge_score_a")
        score_b = genome.metadata.get("merge_score_b")
        if not isinstance(score_a, (int, float)) or isinstance(score_a, bool):
            score_a = genome_score(self.get(genome.parent_id))
        if not isinstance(score_b, (int, float)) or isinstance(score_b, bool):
            context_ids = genome.metadata.get("other_context_ids") or []
            score_b = genome_score(self.get(context_ids[0] if context_ids else None))
        return genome_score(genome) >= max(float(score_a), float(score_b))

    def add(self, genome: Genome) -> bool:
        # Galapagos records invalid normal and merge attempts in the evolutionary trajectory, but
        # neither is an archive member. This is intentionally stricter than the source merge path's
        # score-only acceptance check and keeps the contract uniform across registered scaffolds.
        if genome.metadata.get("valid") is False:
            genome.metadata["gepa_native_admitted"] = False
            genome.metadata.update(admitted=False, eval_failed=True)
            return False

        is_merge = bool(genome.metadata.get("gepa_native_merge"))
        if genome.parent_id is not None:
            if is_merge:
                if not self._merge_is_acceptable(genome):
                    genome.metadata["gepa_native_admitted"] = False
                    genome.metadata["admitted"] = False
                    return False
            else:
                if self.acceptance_gating:
                    parent = self.get(genome.parent_id)
                    parent_score = genome_score(parent) if parent is not None else 0.0
                    if genome_score(genome) <= parent_score:
                        self.add_rejected(genome)
                        return False

        self._add_to_archive(genome)
        genome.metadata.pop("eval_failed", None)
        genome.metadata["gepa_native_admitted"] = True
        genome.metadata["admitted"] = True
        return True

    def _add_to_archive(self, genome: Genome) -> None:
        if not self._programs:
            self.initial_program_id = genome.id
        self._programs[genome.id] = genome

        if genome.id not in self.elite_pool:
            self.elite_pool.append(genome.id)
        self.elite_pool.sort(
            key=lambda program_id: genome_score(self._programs.get(program_id)),
            reverse=True,
        )

        # Preserve SkyDiscover's pinning loop exactly.  Because pinned members do not consume the
        # ``len(keep) < population_size`` allowance consistently, the sampling pool can occasionally
        # exceed the nominal bound; tightening it would change the source algorithm.
        if len(self.elite_pool) > self.population_size:
            pinned = {
                self.best_program_id,
                self.initial_program_id,
                genome.id,
            } - {None}
            keep: list[str] = []
            for program_id in self.elite_pool:
                if program_id in pinned or len(keep) < self.population_size:
                    keep.append(program_id)
            self.elite_pool = keep

        for metric_name, value in (genome.scores or {}).items():
            # Upstream uses this raw isinstance check, so bool-valued metrics participate here even
            # though they are excluded from scalar score fallback.
            if not isinstance(value, (int, float)):
                continue
            current = self.metric_best.get(metric_name)
            if current is None or value > current[1]:
                self.metric_best[metric_name] = (genome.id, value)
                self.program_at_metric_front[metric_name] = {genome.id}
            elif value == current[1]:
                self.program_at_metric_front.setdefault(metric_name, set()).add(genome.id)

        if (
            self.best_program_id is None
            or genome_score(genome) > genome_score(self._programs.get(self.best_program_id))
        ):
            self.best_program_id = genome.id

    def add_rejected(self, genome: Genome) -> None:
        genome.metadata["gepa_native_admitted"] = False
        genome.metadata["admitted"] = False
        self.rejection_history.append(genome)

    def get_rejection_history(self, limit: int | None = None) -> list[Genome]:
        rejected = list(self.rejection_history)
        return rejected[-limit:] if limit is not None else rejected

    def elite_programs(self) -> list[Genome]:
        return [
            self._programs[program_id]
            for program_id in self.elite_pool
            if program_id in self._programs
        ]

    def query(self, spec: dict | None = None) -> list[Genome]:
        spec = spec or {}
        members = self.elite_programs() if spec.get("elite") else list(self._programs.values())
        if spec.get("sorted") or "top" in spec:
            members.sort(key=genome_score, reverse=True)
        top = spec.get("top")
        return members[:top] if top is not None else members

    def all(self) -> list[Genome]:
        """Return the accepted full archive; the elite pool is sampling-only."""
        return list(self._programs.values())

    def best(self) -> Genome | None:
        if self.best_program_id in self._programs:
            return self._programs[self.best_program_id]
        return max(self._programs.values(), key=genome_score) if self._programs else None

    def state_dict(self) -> dict:
        return {
            "elite_pool": list(self.elite_pool),
            "initial_program_id": self.initial_program_id,
            "metric_best": {
                metric: [program_id, value]
                for metric, (program_id, value) in self.metric_best.items()
            },
            "program_at_metric_front": {
                metric: list(program_ids)
                for metric, program_ids in self.program_at_metric_front.items()
            },
            "rejection_history": [
                genome_to_dict(genome) for genome in self.rejection_history
            ],
        }

    def load_state_dict(self, state: dict) -> None:
        if not isinstance(state, dict):
            return
        if not state:
            self._rebuild_native_indexes()
            return

        self.elite_pool = [
            program_id
            for program_id in state.get("elite_pool", [])
            if program_id in self._programs
        ]
        initial = state.get("initial_program_id")
        self.initial_program_id = initial if initial in self._programs else self.initial_program_id

        self.metric_best.clear()
        for metric, value in (state.get("metric_best") or {}).items():
            if (
                isinstance(value, (list, tuple))
                and len(value) >= 2
                and value[0] in self._programs
            ):
                self.metric_best[str(metric)] = (value[0], value[1])

        self.program_at_metric_front = {
            str(metric): {
                program_id for program_id in (program_ids or [])
                if program_id in self._programs
            }
            for metric, program_ids in (state.get("program_at_metric_front") or {}).items()
        }

        self.rejection_history.clear()
        for raw in state.get("rejection_history", []):
            if not isinstance(raw, dict):
                continue
            try:
                self.rejection_history.append(genome_from_dict(raw))
            except Exception:
                continue

        if not self.elite_pool and self._programs:
            self._rebuild_native_indexes()

    def _rebuild_native_indexes(self) -> None:
        """SkyDiscover legacy-checkpoint fallback: derive native indexes from the archive."""
        self.elite_pool = sorted(
            self._programs,
            key=lambda program_id: genome_score(self._programs[program_id]),
            reverse=True,
        )[: self.population_size]
        if self._programs:
            self.initial_program_id = min(
                self._programs,
                key=lambda program_id: int(
                    self._programs[program_id].metadata.get("iteration", 0) or 0
                ),
            )
            self.best_program_id = max(
                self._programs,
                key=lambda program_id: genome_score(self._programs[program_id]),
            )

        self.metric_best.clear()
        self.program_at_metric_front.clear()
        for program_id, genome in self._programs.items():
            for metric_name, value in (genome.scores or {}).items():
                if not isinstance(value, (int, float)):
                    continue
                current = self.metric_best.get(metric_name)
                if current is None or value > current[1]:
                    self.metric_best[metric_name] = (program_id, value)
                    self.program_at_metric_front[metric_name] = {program_id}
                elif value == current[1]:
                    self.program_at_metric_front.setdefault(metric_name, set()).add(program_id)
