"""SkyDiscover ``gepa_native`` registered as an isolated Galapagos scaffold."""
from __future__ import annotations

import logging
import time

from ...components.proposer import Env
from ...config import GalapagosConfig
from ...models import GalapagosModel
from ...models.base import Prompt
from ...records import Selection
from ..base_scaffold import GalapagosScaffold
from ..registry import register_scaffold
from .memory import GepaNativeMemory
from .population import GepaNativePopulation, genome_score
from .prompt_builder import GepaNativePromptBuilder
from .proposer import GepaNativeMergeProposer, GepaNativeProposer
from .selection_policy import GepaNativeSelectionPolicy

log = logging.getLogger(__name__)


@register_scaffold("gepa_native")
class GepaNativeScaffold(GalapagosScaffold):
    """Stored-parent guided evolution with rejection reflection and LLM crossover."""

    name = "gepa_native"

    def __init__(self, **kwargs):
        super().__init__(**kwargs)
        self.use_merge = bool(self.config.proposer.use_merge)
        self.merge_after_stagnation = int(self.config.general.merge_after_stagnation)
        self.max_merge_attempts = int(self.config.proposer.max_merge_attempts)
        self._best_score_seen = float("-inf")
        self._iterations_without_improvement = 0
        self._merge_due = False
        self._merge_attempts_used = 0
        self._merge_pairs_tried: set[tuple[str, str]] = set()
        self._merge_proposer = GepaNativeMergeProposer()

    @classmethod
    def build_components(
        cls,
        config: GalapagosConfig,
        model: GalapagosModel | None,
    ) -> dict:
        population_config = config.population
        selection_config = config.selection_policy
        prompt_config = config.prompt_builder
        population = GepaNativePopulation(
            population_size=int(population_config.population_size),
            max_rejection_history=int(population_config.max_rejection_history),
            acceptance_gating=bool(population_config.acceptance_gating),
        )
        return {
            "population": population,
            "selection_policy": GepaNativeSelectionPolicy(
                seed=int(config.seed or 42),
                candidate_selection_strategy=selection_config.candidate_selection_strategy,
                epsilon=float(selection_config.epsilon),
                num_inspirations=int(selection_config.num_inspirations),
            ),
            "prompt_builder": GepaNativePromptBuilder(
                population=population,
                max_recent_failures=int(prompt_config.max_recent_failures),
            ),
            "proposer": GepaNativeProposer(),
            "memory": GepaNativeMemory(),
        }

    def setup(self, task) -> None:
        super().setup(task)
        best = self.population.best()
        seed_feedback = getattr(
            getattr(self, "_seed_eval_result", None),
            "text_feedback",
            None,
        )
        if best is not None and seed_feedback:
            # Base setup preserves the EvalResult envelope specifically for methods whose first
            # proposal needs seed-only evidence. SkyDiscover renders evaluator diagnostics from the
            # seed on iteration one, before after_step has ever run.
            best.artifacts.setdefault("text_feedback", seed_feedback)
        self._best_score_seen = genome_score(best)

    def step(self, task) -> None:
        """Keep SkyDiscover's iteration-level fail-soft controller behavior."""
        self._select_s = self._prompt_s = self._llm_s = self._eval_s = None
        try:
            super().step(task)
        except Exception as exc:  # noqa: BLE001 - upstream converts the whole iteration to result.error
            self._llm_total += self._llm_s or 0.0
            self._eval_total += self._eval_s or 0.0
            self._select_total += self._select_s or 0.0
            self._prompt_total += self._prompt_s or 0.0
            self._iterations_without_improvement += 1
            log.warning(
                "gepa_native iteration %s failed and was skipped: %s",
                self.ctx.iteration,
                exc,
                exc_info=True,
            )

    def before_step(self) -> None:
        # ``GalapagosScaffold.step`` has already incremented ctx.iteration here.  The source loop also
        # performs a scheduled merge at the beginning of the numbered iteration.
        if self._merge_due:
            self._merge_due = False
            self._attempt_merge(self.ctx.iteration)

    def after_step(self, child, result) -> None:
        text_feedback = getattr(result, "text_feedback", None) if result is not None else None
        if text_feedback:
            child.artifacts.setdefault("text_feedback", text_feedback)

        # The shared Genome.fitness uses the same numeric-mean fallback for ordinary task results, but
        # SkyDiscover defines an empty/non-numeric verdict as 0 rather than -inf.  Keep the controller's
        # winner synchronized to the native population's score helper for that edge case too.
        native_best = self.population.best()
        if native_best is not None:
            self.ctx.best = native_best

        if child.metadata.get("gepa_native_merge"):
            if child.metadata.get("gepa_native_admitted"):
                child_score = genome_score(child)
                if child_score > self._best_score_seen:
                    self._best_score_seen = child_score
                # Source resets stagnation after every accepted merge, including an equal-score merge.
                self._iterations_without_improvement = 0
            # A rejected/failed merge leaves the counter unchanged.
            return

        admitted = bool(child.metadata.get("gepa_native_admitted"))
        if result is None or not bool(result.valid) or not admitted:
            self._iterations_without_improvement += 1
            return

        if self.use_merge and self._merge_attempts_used < self.max_merge_attempts:
            self._merge_due = True

        child_score = genome_score(child)
        if child_score > self._best_score_seen:
            self._best_score_seen = child_score
            self._iterations_without_improvement = 0
        else:
            # Improving over a non-best parent is accepted but still counts as global stagnation.
            self._iterations_without_improvement += 1

    def periodic(self) -> None:
        if self._should_merge():
            self._attempt_merge(self.ctx.iteration)

    def _should_merge(self) -> bool:
        return (
            self.use_merge
            and self._iterations_without_improvement >= self.merge_after_stagnation
            and self._merge_attempts_used < self.max_merge_attempts
        )

    def _attempt_merge(self, iteration: int) -> None:
        if self._merge_attempts_used >= self.max_merge_attempts:
            return
        if len(self.population.programs) < 2:
            return

        program_a, program_b = self.selection_policy.merge_candidates(self.population)
        if program_a is None or program_b is None or program_a.id == program_b.id:
            return

        pair_key = tuple(sorted((program_a.id, program_b.id)))
        if pair_key in self._merge_pairs_tried:
            return
        # Budget is consumed before model generation, exactly as in SkyDiscover.
        self._merge_pairs_tried.add(pair_key)
        self._merge_attempts_used += 1

        score_a = genome_score(program_a)
        score_b = genome_score(program_b)
        prompt = self._build_merge_prompt(program_a, program_b)
        selection = Selection(
            parent=program_a,
            inspirations=[program_b],
            pool=self.population.all(),
            details={
                "selection_strategy": "gepa_native_merge",
                "selection_mode": "complementary_pair",
                "merge_pair": list(pair_key),
                "merge_attempt": self._merge_attempts_used,
            },
        )
        env = Env(
            model=self.model,
            selection=selection,
            evaluator=self.evaluator,
            memory=self.memory,
            ctx=self.ctx,
            attempt_number=1,
        )

        self._select_s = self._prompt_s = None
        started = time.monotonic()
        try:
            child = self._merge_proposer.propose(prompt, env)
        except Exception as exc:  # noqa: BLE001 - merge failures never abort normal discovery
            self._llm_s = time.monotonic() - started
            self._llm_total += self._llm_s
            log.warning("gepa_native merge model call failed: %s", exc)
            return
        self._llm_s = time.monotonic() - started
        self._llm_total += self._llm_s
        model_call = child.metadata.get("model_call") or {}
        self._cached_tokens += int(model_call.get("cached_tokens", 0) or 0)
        child.trace.setdefault(
            "proposal_prompt",
            {"system": prompt.system, "user": prompt.user},
        )
        child.trace.setdefault("proposal_duration_seconds", self._llm_s)
        child.metadata.update({
            "gepa_native_merge": True,
            "changes": "LLM-mediated merge",
            "merge_score_a": score_a,
            "merge_score_b": score_b,
            "parent_metrics": dict(program_a.scores),
            "other_context_ids": [program_b.id],
            "parent_info": ["Merge Parent A", program_a.id],
            "context_info": [["Merge Parent B", program_b.id]],
            "iteration": iteration,
        })

        if not child.content:
            self._record_with_merge_proposer(
                child,
                None,
                prompt=prompt,
                selection=selection,
                admitted=False,
                reason="no_diff",
                eval_source=None,
            )
            return

        started = time.monotonic()
        try:
            result = self._evaluate_search_candidate(child)
        except Exception as exc:  # noqa: BLE001 - merge evaluator exceptions are fail-soft upstream
            self._eval_s = time.monotonic() - started
            self._eval_total += self._eval_s
            child.trace["evaluation_duration_seconds"] = self._eval_s
            self._record_with_merge_proposer(
                child,
                None,
                prompt=prompt,
                selection=selection,
                admitted=None,
                reason="evaluation_failed",
                evaluation_error=exc,
            )
            log.warning("gepa_native merge evaluation failed: %s", exc)
            return

        self._eval_s = time.monotonic() - started
        self._eval_total += self._eval_s
        child.trace["evaluation_duration_seconds"] = self._eval_s
        child.scores = result.metrics
        child.artifacts.update(result.artifacts)
        if getattr(result, "text_feedback", None):
            child.artifacts.setdefault("text_feedback", result.text_feedback)
        child.metadata["valid"] = bool(result.valid)
        if not result.valid:
            child.metadata["gepa_native_admitted"] = False
        child.metadata["inner_attempts_used"] = 1

        # The source merge path asks only whether the scalar score meets both parents. Galapagos first
        # applies its framework-wide evaluator-validity contract, then the native score merge gate.
        original_proposer = self.proposer
        self.proposer = self._merge_proposer
        try:
            self._complete_attempt(
                child,
                result,
                "ok",
                prompt=prompt,
                selection=selection,
            )
        finally:
            self.proposer = original_proposer

    def _record_with_merge_proposer(
        self,
        child,
        result,
        *,
        prompt,
        selection,
        admitted,
        reason,
        eval_source="evaluator",
        evaluation_error=None,
    ) -> None:
        """Record an auxiliary merge with crossover semantics and both ETIF parents."""
        original_proposer = self.proposer
        self.proposer = self._merge_proposer
        try:
            self._record(
                child,
                result,
                origin="proposal",
                admitted=admitted,
                reason=reason,
                eval_source=eval_source,
                prompt=prompt,
                selection=selection,
                attempt=1,
                evaluation_error=evaluation_error,
            )
        finally:
            self.proposer = original_proposer

    @staticmethod
    def _build_merge_prompt(program_a, program_b) -> Prompt:
        score_a = genome_score(program_a)
        score_b = genome_score(program_b)

        strengths_a: list[str] = []
        strengths_b: list[str] = []
        if program_a.scores and program_b.scores:
            # Deliberately retain the source set iteration order; sorting changes prompt bytes.
            for key in set(program_a.scores) | set(program_b.scores):
                value_a = program_a.scores.get(key)
                value_b = program_b.scores.get(key)
                if (
                    not isinstance(value_a, (int, float))
                    or not isinstance(value_b, (int, float))
                ):
                    continue
                if value_a > value_b:
                    strengths_a.append(f"{key}: {value_a}")
                elif value_b > value_a:
                    strengths_b.append(f"{key}: {value_b}")

        strengths = ""
        if strengths_a or strengths_b:
            strengths = "\n## Per-Metric Strengths\n"
            if strengths_a:
                strengths += f"Program A leads on: {', '.join(strengths_a)}\n"
            if strengths_b:
                strengths += f"Program B leads on: {', '.join(strengths_b)}\n"

        diagnostics = ""
        for label, program in (("A", program_a), ("B", program_b)):
            parts = []
            for key, value in (program.artifacts or {}).items():
                # Galapagos keeps the raw proposal alongside evaluator artifacts; SkyDiscover does
                # not.  Excluding those two framework-only keys preserves the source prompt content.
                if key in {"response", "reasoning"}:
                    continue
                if not isinstance(value, str) or not value.strip():
                    continue
                display = value if len(value) <= 500 else value[:500] + "... (truncated)"
                parts.append(f"- {key}: {display}")
            if parts:
                diagnostics += (
                    f"\n## Program {label} Diagnostics\n"
                    + "\n".join(parts)
                    + "\n"
                )

        system = (
            "You are an expert programmer. Your task is to merge two programs into "
            "a single improved program that combines the strengths of both. "
            "Output only the complete merged program inside a code block."
        )
        user = (
            f"## Program A (score: {score_a:.4f})\n"
            f"```\n{program_a.content}\n```\n\n"
            f"## Program B (score: {score_b:.4f})\n"
            f"```\n{program_b.content}\n```\n"
            f"{strengths}"
            f"{diagnostics}\n"
            "## Instructions\n"
            "Combine the best ideas from both programs into a single solution. "
            "Preserve any approach that contributes to a higher score. "
            "Resolve conflicts by choosing the strategy that is more likely to "
            "generalise across all test cases. Output the complete merged program."
        )
        return Prompt(system=system, user=user)

    def _checkpoint_scaffold_state(self) -> dict:
        """Persist controller state so a Galapagos resume matches an uninterrupted native run."""
        return {
            "best_score_seen": self._best_score_seen,
            "iterations_without_improvement": self._iterations_without_improvement,
            "merge_due": self._merge_due,
            "merge_attempts_used": self._merge_attempts_used,
            "merge_pairs_tried": [list(pair) for pair in sorted(self._merge_pairs_tried)],
        }

    def _load_checkpoint_scaffold_state(self, state: dict) -> None:
        state = state if isinstance(state, dict) else {}
        best = self.population.best()
        if best is not None:
            self.ctx.best = best
        self._best_score_seen = float(
            state.get("best_score_seen", genome_score(best))
        )
        self._iterations_without_improvement = int(
            state.get("iterations_without_improvement", 0)
        )
        self._merge_due = bool(state.get("merge_due", False))
        self._merge_attempts_used = int(state.get("merge_attempts_used", 0))
        self._merge_pairs_tried = {
            tuple(sorted((str(pair[0]), str(pair[1]))))
            for pair in state.get("merge_pairs_tried", [])
            if isinstance(pair, (list, tuple)) and len(pair) == 2
        }
