"""RelationSyncService — entity_relation 체크포인트 → RelationFact/RelationParticipant.

Phase 4.4: 기존 "visual_variant 전체 DELETE 후 재생성" 경로를 delta sync
(UPSERT)로 전환. `(base_canon_id, variant_canon_id)` 논리 키 기준으로 동일한
관계는 UPDATE(continuity_reason만 갱신), 사라진 관계만 DELETE, 새 관계만 INSERT.
"""
from __future__ import annotations

import uuid
from typing import Dict, Tuple

from sqlalchemy import text as sql_text

from app.services.checkpoint_sync._base import BaseSyncService, is_cp_syncable


class RelationSyncService(BaseSyncService):
    def sync_from_checkpoint(self) -> Dict[str, int]:
        """entity_relation 체크포인트 → RelationFact + RelationParticipant (delta sync).

        Returns: {"relations": n_current, "inserted": n, "updated": n, "deleted": n, "skipped": n}
        """
        from app.models.project import EntityCanon, RelationFact, RelationParticipant

        rel_cp = self._load_cp("entity_relation")
        # M2 Fix 2 단일 표준: partial이라도 데이터 있으면 sync.
        # 빈 relations(cascade 직후)면 is_cp_syncable이 False 반환 — 별도 검사 불필요.
        if not is_cp_syncable(rel_cp, ["relations"]):
            return self._empty_delta(skipped=1)

        # Codex P1-2: partial 시 stale 제거 skip — 부분 cp가 미포함 relation을
        # 정상 누락으로 처리해 옛 row 삭제하지 않도록. completed만 authoritative.
        is_partial = (rel_cp.get("status") == "partial")

        # short_id → entity_canon.id 매핑
        canon_rows = (
            self.db.query(EntityCanon)
            .filter(
                EntityCanon.project_id == self.project_id,
                EntityCanon.entity_type.in_(["character", "location", "prop"]),
            )
            .all()
        )
        sid_to_canon_id = {c.short_id: c.id for c in canon_rows}

        # 기대 상태 (체크포인트에서 parse)
        # key: (base_canon_id, variant_canon_id) → reason
        # P2-2: completed cp는 data.relations 키 자체가 없을 수도 있음 (빈 zero rows).
        desired: Dict[Tuple[str, str], str] = {}
        for rel in rel_cp.get("data", {}).get("relations", []) or []:
            if not rel.get("visual_similarity"):
                continue
            base_cid = sid_to_canon_id.get(rel.get("base_short_id", ""))
            var_cid = sid_to_canon_id.get(rel.get("variant_short_id", ""))
            if not base_cid or not var_cid:
                self.logger.warning(
                    "Relation skip — unknown canon: %s→%s",
                    rel.get("base_short_id"),
                    rel.get("variant_short_id"),
                )
                continue
            desired[(base_cid, var_cid)] = rel.get(
                "reason", "시각적 변형 — 기본 요소에 의존"
            )

        # 현재 상태 (DB의 visual_variant)
        # key: (base_canon_id, variant_canon_id) → (rel_id, continuity_reason)
        existing = self._load_existing_visual_variants()

        # Delta 계산
        desired_keys = set(desired.keys())
        existing_keys = set(existing.keys())

        to_insert = desired_keys - existing_keys
        # Codex P1-2: partial이면 stale relation 삭제 skip — 부분 cp가 옛 relation을
        # 정상 누락으로 처리해 데이터 손실되지 않도록 (다음 completed에서 정리).
        to_delete = set() if is_partial else (existing_keys - desired_keys)
        to_update = [
            key for key in desired_keys & existing_keys
            if desired[key] != existing[key][1]
        ]

        # 삭제 (participant 먼저 → fact)
        for key in to_delete:
            rel_id = existing[key][0]
            self.db.execute(
                sql_text("DELETE FROM relation_participant WHERE relation_id = :rid"),
                {"rid": rel_id},
            )
            self.db.execute(
                sql_text("DELETE FROM relation_fact WHERE id = :rid"),
                {"rid": rel_id},
            )

        # 업데이트 (continuity_reason만)
        for key in to_update:
            rel_id = existing[key][0]
            self.db.execute(
                sql_text(
                    "UPDATE relation_fact SET continuity_reason = :reason WHERE id = :rid"
                ),
                {"reason": desired[key], "rid": rel_id},
            )

        # 신규 (RelationFact + 2 participant)
        for key in to_insert:
            base_cid, var_cid = key
            rel_id = str(uuid.uuid4())
            self.db.add(
                RelationFact(
                    id=rel_id,
                    project_id=self.project_id,
                    relation_family="identity",
                    relation_type="visual_variant",
                    directionality="directed",
                    temporal_scope="persistent",
                    continuity_priority="critical",
                    continuity_reason=desired[key],
                    created_at=self.now,
                )
            )
            self.db.add(
                RelationParticipant(
                    id=str(uuid.uuid4()),
                    relation_id=rel_id,
                    canon_id=base_cid,
                    participant_role="base",
                    participant_order=1,
                )
            )
            self.db.add(
                RelationParticipant(
                    id=str(uuid.uuid4()),
                    relation_id=rel_id,
                    canon_id=var_cid,
                    participant_role="variant",
                    participant_order=2,
                )
            )

        if to_insert or to_update or to_delete:
            self.db.flush()
            self.logger.info(
                "Visual-variant delta sync: +%d / ~%d / -%d (total=%d)",
                len(to_insert),
                len(to_update),
                len(to_delete),
                len(desired),
            )

        return {
            "relations": len(desired),
            "inserted": len(to_insert),
            "updated": len(to_update),
            "deleted": len(to_delete),
            "skipped": 0,
        }

    def _load_existing_visual_variants(self) -> Dict[Tuple[str, str], Tuple[str, str]]:
        """현재 DB의 visual_variant 관계를 (base_canon_id, variant_canon_id) → (rel_id, reason) 맵으로.

        RelationParticipant를 RelationFact와 JOIN하여 base/variant role을 구분.
        RelationFact가 비정상적으로 participant를 2개 미만 가진 경우 skip.
        """
        rows = self.db.execute(
            sql_text(
                "SELECT rf.id AS rel_id, "
                "       rf.continuity_reason AS reason, "
                "       MAX(CASE WHEN rp.participant_role = 'base' THEN rp.canon_id END) AS base_cid, "
                "       MAX(CASE WHEN rp.participant_role = 'variant' THEN rp.canon_id END) AS variant_cid "
                "FROM relation_fact rf "
                "LEFT JOIN relation_participant rp ON rp.relation_id = rf.id "
                "WHERE rf.project_id = :pid "
                "  AND rf.relation_type = 'visual_variant' "
                "GROUP BY rf.id, rf.continuity_reason"
            ),
            {"pid": self.project_id},
        ).fetchall()

        result: Dict[Tuple[str, str], Tuple[str, str]] = {}
        for row in rows:
            rel_id, reason, base_cid, variant_cid = row[0], row[1], row[2], row[3]
            if not base_cid or not variant_cid:
                # 비정상 (participant 누락) — 정리 대상이지만 delta에서 skip
                self.logger.warning(
                    "visual_variant relation %s has missing base/variant, skipping delta",
                    rel_id,
                )
                continue
            result[(base_cid, variant_cid)] = (rel_id, reason or "")
        return result

    @staticmethod
    def _empty_delta(*, skipped: int = 0) -> Dict[str, int]:
        return {
            "relations": 0,
            "inserted": 0,
            "updated": 0,
            "deleted": 0,
            "skipped": skipped,
        }
