from __future__ import annotations

from copy import deepcopy
from enum import Enum
from typing import Any

from .compiler import compile_solution
from .normalizer import normalize_solution_spec
from .spec_models import SolutionSpec


class DesignStage(str, Enum):
    discover = "discover"
    design = "design"
    experience = "experience"
    finalize = "finalize"


STAGE_ORDER = [DesignStage.discover, DesignStage.design, DesignStage.experience]

LIST_MERGE_KEYS = {
    "entities": "entity_id",
    "roles": "role_id",
    "requirements": "requirement_id",
    "success_metrics": "metric_id",
    "fields": "field_id",
    "subfields": "field_id",
    "relations": "relation_id",
    "lifecycle_stages": "stage_id",
    "views": "view_id",
    "charts": "chart_id",
    "nodes": "node_id",
    "sections": "section_id",
    "items": "item_id",
    "children": "item_id",
    "entity_scopes": "entity_id",
}


def merge_design_payload(base: Any, patch: Any, *, key_hint: str | None = None) -> Any:
    if patch is None:
        return deepcopy(base)
    if base is None:
        return deepcopy(patch)
    if isinstance(base, dict) and isinstance(patch, dict):
        merged = deepcopy(base)
        for key, value in patch.items():
            merged[key] = merge_design_payload(merged.get(key), value, key_hint=key)
        return merged
    if isinstance(base, list) and isinstance(patch, list):
        merge_key = LIST_MERGE_KEYS.get(key_hint or "")
        if merge_key and _can_merge_named_list(base, patch, merge_key):
            return _merge_named_list(base, patch, merge_key)
        return deepcopy(patch)
    return deepcopy(patch)


def evaluate_design_session(stage_payloads: dict[str, dict[str, Any]]) -> dict[str, Any]:
    discover_spec = merge_design_payload({}, stage_payloads.get(DesignStage.discover.value, {}))
    design_spec = merge_design_payload(discover_spec, stage_payloads.get(DesignStage.design.value, {}))
    experience_spec = merge_design_payload(design_spec, stage_payloads.get(DesignStage.experience.value, {}))

    discover_missing = _validate_discover_stage(discover_spec)
    design_missing = _blocked("discover", discover_missing) if discover_missing else _validate_design_stage(design_spec)
    experience_missing = _blocked("design", design_missing) if design_missing else _validate_experience_stage(experience_spec)

    stage_results = {
        DesignStage.discover.value: _stage_result(discover_missing),
        DesignStage.design.value: _stage_result(design_missing, blocked=bool(discover_missing)),
        DesignStage.experience.value: _stage_result(experience_missing, blocked=bool(design_missing)),
    }
    current_stage = _current_stage(discover_missing, design_missing, experience_missing)
    status = "ready" if current_stage == DesignStage.finalize.value else "active"
    return {
        "status": status,
        "current_stage": current_stage,
        "next_stage": None if current_stage == DesignStage.finalize.value else current_stage,
        "stage_results": stage_results,
        "merged_design_spec": experience_spec,
    }


def finalize_design_session(stage_payloads: dict[str, dict[str, Any]]) -> dict[str, Any]:
    evaluation = evaluate_design_session(stage_payloads)
    if evaluation["current_stage"] != DesignStage.finalize.value:
        raise ValueError("design session is not ready to finalize")
    parsed = SolutionSpec.model_validate(evaluation["merged_design_spec"])
    normalized = normalize_solution_spec(parsed)
    compiled = compile_solution(normalized)
    return {
        **evaluation,
        "normalized_solution_spec": normalized.model_dump(mode="json"),
        "execution_plan": compiled.execution_plan.as_dict(),
    }


def _stage_result(missing_requirements: list[str], *, blocked: bool = False) -> dict[str, Any]:
    if blocked and missing_requirements:
        status = "blocked"
    elif missing_requirements:
        status = "pending"
    else:
        status = "completed"
    return {
        "status": status,
        "missing_requirements": missing_requirements,
    }


def _current_stage(discover_missing: list[str], design_missing: list[str], experience_missing: list[str]) -> str:
    if discover_missing:
        return DesignStage.discover.value
    if design_missing:
        return DesignStage.design.value
    if experience_missing:
        return DesignStage.experience.value
    return DesignStage.finalize.value


def _blocked(previous_stage: str, previous_missing: list[str]) -> list[str]:
    return [f"{previous_stage} stage is incomplete"] if previous_missing else []


def _validate_discover_stage(spec: dict[str, Any]) -> list[str]:
    missing: list[str] = []
    solution_name = spec.get("solution_name")
    if not isinstance(solution_name, str) or not solution_name.strip():
        missing.append("solution_name is required")

    entities = spec.get("entities")
    if not isinstance(entities, list) or not entities:
        missing.append("entities must be a non-empty list")
        return missing

    seen_entity_ids: set[str] = set()
    for entity in entities:
        entity_id = entity.get("entity_id")
        display_name = entity.get("display_name")
        kind = entity.get("kind")
        if not entity_id:
            missing.append("each entity must declare entity_id")
            continue
        if entity_id in seen_entity_ids:
            missing.append(f"entity '{entity_id}' is duplicated")
            continue
        seen_entity_ids.add(entity_id)
        if not display_name:
            missing.append(f"entity '{entity_id}' must declare display_name")
        if not kind:
            missing.append(f"entity '{entity_id}' must declare kind")
        fields = entity.get("fields")
        if not isinstance(fields, list) or not fields:
            missing.append(f"entity '{entity_id}' must declare fields")
            continue
        field_ids = {field.get("field_id") for field in fields if isinstance(field, dict)}
        if kind in {"master", "transaction"}:
            title_field_id = entity.get("title_field_id")
            if not title_field_id:
                missing.append(f"entity '{entity_id}' must explicitly declare title_field_id")
            elif title_field_id not in field_ids:
                missing.append(f"entity '{entity_id}' title_field_id '{title_field_id}' is missing from fields")
        if entity.get("workflow") or entity.get("lifecycle_stages"):
            status_field_id = entity.get("status_field_id")
            if not status_field_id:
                missing.append(f"entity '{entity_id}' must explicitly declare status_field_id")
            elif status_field_id not in field_ids:
                missing.append(f"entity '{entity_id}' status_field_id '{status_field_id}' is missing from fields")
    return missing


def _validate_design_stage(spec: dict[str, Any]) -> list[str]:
    missing: list[str] = []
    for entity in spec.get("entities", []):
        entity_id = entity.get("entity_id", "<unknown>")
        if "form_layout" not in entity:
            missing.append(f"entity '{entity_id}' must explicitly declare form_layout")
        if "workflow" not in entity:
            missing.append(f"entity '{entity_id}' must explicitly declare workflow")
        if "views" not in entity:
            missing.append(f"entity '{entity_id}' must explicitly declare views")
        if "charts" not in entity:
            missing.append(f"entity '{entity_id}' must explicitly declare charts")
        if "sample_records" not in entity:
            missing.append(f"entity '{entity_id}' must explicitly declare sample_records")
    return missing


def _validate_experience_stage(spec: dict[str, Any]) -> list[str]:
    missing: list[str] = []
    if "portal" not in spec:
        missing.append("portal must be explicitly declared")
    else:
        portal = spec.get("portal") or {}
        if portal.get("enabled", True) and not portal.get("sections"):
            missing.append("portal.sections must be provided when portal.enabled is true")
    if "navigation" not in spec:
        missing.append("navigation must be explicitly declared")
    else:
        navigation = spec.get("navigation") or {}
        if navigation.get("enabled", True) and not navigation.get("items"):
            missing.append("navigation.items must be provided when navigation.enabled is true")
    return missing


def _can_merge_named_list(base: list[Any], patch: list[Any], merge_key: str) -> bool:
    items = [*base, *patch]
    if not items:
        return False
    return all(isinstance(item, dict) and merge_key in item for item in items)


def _merge_named_list(base: list[dict[str, Any]], patch: list[dict[str, Any]], merge_key: str) -> list[dict[str, Any]]:
    merged = [deepcopy(item) for item in base]
    index = {item[merge_key]: position for position, item in enumerate(merged)}
    for item in patch:
        item_key = item[merge_key]
        if item_key in index:
            merged[index[item_key]] = merge_design_payload(merged[index[item_key]], item)
        else:
            index[item_key] = len(merged)
            merged.append(deepcopy(item))
    return merged
