"""
Rollback service: restore Magento products to a previous version.
"""

from __future__ import annotations

import logging
from datetime import datetime, timezone
from typing import Any, Dict, List, Optional, Protocol

from magento.magento_api import MagentoRestClient

logger = logging.getLogger(__name__)


class VersionRepository(Protocol):
    def get_version_by_id(self, version_id: int) -> Optional[dict]: ...
    def get_versions_as_of(
        self, connection_id: int, as_of: datetime, sku: Optional[str] = None
    ) -> List[dict]: ...


class SyncStateRepository(Protocol):
    def upsert_state(
        self,
        connection_id: int,
        sku: str,
        *,
        data_hash: Optional[str] = None,
        images_hash: Optional[str] = None,
        relations_hash: Optional[str] = None,
        last_pushed_at: Optional[datetime] = None,
        last_error: Optional[str] = None,
    ) -> None: ...


class MagentoRollbackService:
    """Restore products to a previous version. Pass 1: product, Pass 2: media, Pass 3: relations."""

    def __init__(
        self,
        api_client: MagentoRestClient,
        connection_id: int,
        version_repo: VersionRepository,
        state_repo: SyncStateRepository,
        *,
        media_executor: Optional[Any] = None,
        relations_executor: Optional[Any] = None,
    ) -> None:
        self._api = api_client
        self._conn_id = connection_id
        self._version_repo = version_repo
        self._state_repo = state_repo
        self._media_executor = media_executor
        self._relations_executor = relations_executor

    def rollback_to_version(
        self,
        *,
        sku: Optional[str] = None,
        all_skus: bool = False,
        version_id: Optional[int] = None,
        as_of: Optional[datetime] = None,
    ) -> Dict[str, Any]:
        """
        Rollback products to a previous version.
        - sku: rollback single SKU only
        - all_skus: rollback every SKU (requires as_of)
        - version_id: use exact version
        - as_of: latest version per SKU with created_at <= as_of (for all_skus or sku)
        """
        versions: List[dict] = []
        if version_id is not None:
            v = self._version_repo.get_version_by_id(version_id)
            if v:
                versions = [v]
            else:
                return {"status": "failed", "error": f"Version {version_id} not found"}
        elif as_of is not None:
            versions = self._version_repo.get_versions_as_of(
                self._conn_id, as_of, sku=sku
            )
            if not versions and not all_skus:
                return {"status": "failed", "error": f"No versions found for sku={sku} as_of={as_of}"}
        else:
            return {"status": "failed", "error": "Provide version_id or as_of"}

        if not versions:
            return {"status": "success", "rolled_back_count": 0, "message": "No versions to rollback"}

        now = datetime.now(timezone.utc)
        product_ok = 0
        product_fail = 0
        media_ok = 0
        relations_ok = 0

        skus_with_product = {v["sku"] for v in versions if v.get("product_payload_json")}
        skus_with_media = {v["sku"] for v in versions if v.get("media_payload_json")}
        skus_with_relations = {v["sku"] for v in versions if v.get("relations_payload_json")}

        version_by_sku = {v["sku"]: v for v in versions}

        # Pass 1: product payloads
        for v in versions:
            payload = v.get("product_payload_json")
            if not payload:
                continue
            sku_val = v["sku"]
            try:
                status, _ = self._api.put_product(sku_val, payload)
                if status in (200, 201):
                    product_ok += 1
                    self._state_repo.upsert_state(
                        self._conn_id, sku_val,
                        data_hash=v.get("data_hash"),
                        last_pushed_at=now,
                        last_error="",
                    )
                else:
                    if status == 404:
                        status2, _ = self._api.post_product(payload)
                        if status2 in (200, 201):
                            product_ok += 1
                            self._state_repo.upsert_state(
                                self._conn_id, sku_val,
                                data_hash=v.get("data_hash"),
                                last_pushed_at=now,
                                last_error="",
                            )
                        else:
                            product_fail += 1
                    else:
                        product_fail += 1
            except Exception as e:
                logger.exception("Rollback product failed for %s", sku_val)
                product_fail += 1
                self._state_repo.upsert_state(
                    self._conn_id, sku_val,
                    last_error=str(e)[:500],
                )

        # Pass 2: media (requires media executor - sync service's execute_upsert_media)
        if self._media_executor and skus_with_media:
            for sku_val in skus_with_media:
                v = version_by_sku.get(sku_val)
                if not v or not v.get("media_payload_json"):
                    continue
                media_list = v["media_payload_json"]
                row = self._media_payload_to_row(media_list, sku_val)
                try:
                    result = self._media_executor.execute_upsert_media(sku_val, row)
                    if result.get("status") == "success":
                        media_ok += 1
                        self._state_repo.upsert_state(
                            self._conn_id, sku_val,
                            images_hash=v.get("images_hash"),
                            last_pushed_at=now,
                        )
                except Exception as e:
                    logger.exception("Rollback media failed for %s", sku_val)

        # Pass 3: relations (configurable parents only)
        if self._relations_executor and skus_with_relations:
            for sku_val in skus_with_relations:
                v = version_by_sku.get(sku_val)
                if not v or not v.get("relations_payload_json"):
                    continue
                rel = v["relations_payload_json"]
                row = self._relations_payload_to_row(rel, sku_val)
                try:
                    result = self._relations_executor.execute_set_relations(sku_val, row)
                    if result.get("status") == "success":
                        relations_ok += 1
                        self._state_repo.upsert_state(
                            self._conn_id, sku_val,
                            relations_hash=v.get("relations_hash"),
                            last_pushed_at=now,
                        )
                except Exception as e:
                    logger.exception("Rollback relations failed for %s", sku_val)

        return {
            "status": "success",
            "product_ok": product_ok,
            "product_fail": product_fail,
            "media_ok": media_ok,
            "relations_ok": relations_ok,
        }

    def _media_payload_to_row(
        self, media_list: List[Dict[str, Any]], sku: str
    ) -> Dict[str, Any]:
        """Convert media_payload_json to row format for execute_upsert_media."""
        urls: List[str] = []
        for m in sorted(media_list, key=lambda x: x.get("position", 0)):
            u = m.get("url", "")
            if u:
                urls.append(u)
        return {
            "sku": sku,
            "name": sku,
            "base_image": urls[0] if urls else "",
            "additional_images": ",".join(urls[1:]) if len(urls) > 1 else "",
        }

    def _relations_payload_to_row(
        self, rel: Dict[str, Any], parent_sku: str
    ) -> Dict[str, Any]:
        """Convert relations_payload_json to row format for execute_set_relations."""
        children = rel.get("children", [])
        attr_codes = rel.get("attr_codes", [])
        child_map = rel.get("child_option_map", {})
        if not children:
            return {"sku": parent_sku, "product_type": "simple"}
        parts: List[str] = []
        for c_sku in children:
            opts = child_map.get(c_sku, {})
            seg = "sku=" + c_sku
            for k, v in opts.items():
                seg += f",{k}={v}"
            parts.append(seg)
        return {
            "sku": parent_sku,
            "product_type": "configurable",
            "variant_list": ",".join(children),
            "configurable_attributes": ",".join(attr_codes),
            "configurable_variations": "|".join(parts),
        }
