#!/usr/bin/env python3
"""Tests for the original-structure Coze migration package."""

from __future__ import annotations

import csv
import json
import re
import unittest
import zipfile
from pathlib import Path


ROOT = Path(__file__).resolve().parents[1]
SOURCE_WORKFLOW = ROOT / "source" / "workflow-draft.json"
CONVERTED = ROOT / "coze" / "converted"
PACKAGE_ROOT = CONVERTED / "Workflow-running_ai_assistant_5_original_structure_v1-draft-0001"
PACKAGE_ZIP = CONVERTED / "Workflow-running_ai_assistant_5_original_structure_v1-draft-0001.zip"
WORKFLOW_YAML = PACKAGE_ROOT / "workflow" / "running_ai_assistant_5_original_structure_v1-draft.yaml"
NODE_MAPPING = CONVERTED / "original-structure-node-mapping.csv"


class OriginalStructurePackageTest(unittest.TestCase):
    @classmethod
    def setUpClass(cls) -> None:
        cls.source = json.loads(SOURCE_WORKFLOW.read_text(encoding="utf-8"))
        cls.yaml_text = WORKFLOW_YAML.read_text(encoding="utf-8")

    def test_package_files_exist(self) -> None:
        self.assertTrue(PACKAGE_ZIP.exists())
        self.assertTrue((PACKAGE_ROOT / "MANIFEST.yml").exists())
        self.assertTrue(WORKFLOW_YAML.exists())
        self.assertTrue(NODE_MAPPING.exists())

    def test_zip_has_coze_import_layout(self) -> None:
        with zipfile.ZipFile(PACKAGE_ZIP) as zf:
            names = zf.namelist()
        self.assertIn("Workflow-running_ai_assistant_5_original_structure_v1-draft-0001/MANIFEST.yml", names)
        self.assertIn(
            "Workflow-running_ai_assistant_5_original_structure_v1-draft-0001/workflow/running_ai_assistant_5_original_structure_v1-draft.yaml",
            names,
        )

    def test_preserves_original_node_and_edge_counts(self) -> None:
        source_nodes = self.source["graph"]["nodes"]
        source_edges = self.source["graph"]["edges"]
        answer_nodes = [node for node in source_nodes if (node.get("data") or {}).get("type") == "answer"]
        yaml_nodes = re.findall(r"^    - id:", self.yaml_text, re.MULTILINE)
        yaml_edges = re.findall(r"^    - source_node:", self.yaml_text, re.MULTILINE)
        self.assertEqual(len(yaml_nodes), len(source_nodes) + 2)
        self.assertEqual(len(yaml_edges), len(source_edges) + len(answer_nodes) + 1)
        self.assertIn("id: \"899999\"", self.yaml_text)
        self.assertIn("id: \"900001\"", self.yaml_text)

    def test_not_replay_or_old_compatibility_package(self) -> None:
        self.assertNotIn("88条Dify基准答案回放", self.yaml_text)
        self.assertNotIn("NOT_DIFY_BUSINESS_OUTPUT", self.yaml_text)
        self.assertNotIn("Dify/Running compatibility package imported", self.yaml_text)
        self.assertIn("Coze 只允许一个固定结束节点", self.yaml_text)

    def test_preserves_branch_ports(self) -> None:
        source_branch_edges = [
            edge for edge in self.source["graph"]["edges"] if edge.get("sourceHandle") and edge.get("sourceHandle") != "source"
        ]
        self.assertEqual(self.yaml_text.count("source_port:"), len(source_branch_edges))

    def test_node_mapping_covers_every_dify_node(self) -> None:
        with NODE_MAPPING.open("r", encoding="utf-8-sig", newline="") as handle:
            rows = list(csv.DictReader(handle))
        self.assertEqual(len(rows), len(self.source["graph"]["nodes"]))
        mapped_ids = {row["dify_node_id"] for row in rows}
        source_ids = {str(node["id"]) for node in self.source["graph"]["nodes"]}
        self.assertEqual(mapped_ids, source_ids)

    def test_expected_runtime_target_types_are_recorded(self) -> None:
        with NODE_MAPPING.open("r", encoding="utf-8-sig", newline="") as handle:
            targets = {row["dify_type"]: row["runtime_target_type"] for row in csv.DictReader(handle)}
        self.assertEqual(targets["llm"], "llm")
        self.assertEqual(targets["knowledge-retrieval"], "knowledge")
        self.assertEqual(targets["if-else"], "condition_pending_native_rebind")
        self.assertEqual(targets["http-request"], "http_pending_native_rebind")
        self.assertEqual(targets["answer"], "answer_template_code_with_shared_end")


if __name__ == "__main__":
    unittest.main()
