from __future__ import annotations

import json
import re
import unittest
from pathlib import Path


ROOT = Path(__file__).resolve().parents[1]
RAW = ROOT / "raw"


def load(name: str):
    return json.loads((RAW / name).read_text(encoding="utf-8"))


WORKFLOW = load("workflow-draft.json")
APP = load("app-detail.json")
SYSTEM_VARIABLES = load("draft-system-variables.json")["items"]
CONVERSATION_VARIABLES = load("draft-conversation-variables.json")["items"]
NODE_VARIABLES = load("draft-variables-page-1.json")["items"] + load("draft-variables-page-2.json")["items"]
NODES = WORKFLOW["graph"]["nodes"]
EDGES = WORKFLOW["graph"]["edges"]
NODE_BY_ID = {node["id"]: node for node in NODES}
NODE_OUTPUTS = {
    node["id"]: set((node.get("data", {}).get("outputs") or {}).keys())
    for node in NODES
}
KNOWN_SELECTORS = {("sys", item["name"]) for item in SYSTEM_VARIABLES}
KNOWN_SELECTORS |= {("conversation", item["name"]) for item in CONVERSATION_VARIABLES}
KNOWN_SELECTORS |= {tuple(item["selector"]) for item in NODE_VARIABLES}
KNOWN_SELECTORS |= {("env", item["name"]) for item in WORKFLOW.get("environment_variables", [])}
IMPLICIT_OUTPUTS_BY_TYPE = {
    "llm": {"text", "usage", "finish_reason"},
    "knowledge-retrieval": {"result"},
    "http-request": {"body", "status_code", "headers"},
    "answer": {"answer", "files"},
    "if-else": {"result", "selected_case_id"},
}


def selectors_in_text(text: str):
    return re.findall(r"\{\{#([^#]+)#\}\}", text or "")


def safe_id(value: str) -> str:
    return re.sub(r"[^0-9A-Za-z_]+", "_", value).strip("_")[:120] or "case"


class WorkflowContractTests(unittest.TestCase):
    def assert_selector(self, selector):
        selector = tuple(selector)
        self.assertTrue(selector in KNOWN_SELECTORS or selector[0] in NODE_BY_ID, selector)
        if selector[0] in NODE_BY_ID and len(selector) > 1:
            known_names = {item["name"] for item in NODE_VARIABLES if item["selector"][0] == selector[0]}
            known_names |= NODE_OUTPUTS.get(selector[0], set())
            known_names |= IMPLICIT_OUTPUTS_BY_TYPE.get(NODE_BY_ID[selector[0]]["data"]["type"], set())
            self.assertIn(selector[1], known_names, f"unknown node output selector: {selector}")

    def test_000_app_identity(self):
        self.assertEqual(APP["id"], "897352d5-faea-4467-9392-8cf70323c764")
        self.assertEqual(APP["name"], "跑步AI助手-5.0")
        self.assertEqual(APP["mode"], "advanced-chat")

    def test_001_workflow_shape(self):
        self.assertEqual(WORKFLOW["id"], "5f5637de-d434-4f3b-a2b1-b569afea3c89")
        self.assertEqual(len(NODES), 58)
        self.assertEqual(len(EDGES), 72)

    def test_002_export_dsl_exists(self):
        text = (RAW / "dify-export-no-secret.yml").read_text(encoding="utf-8")
        self.assertIn("name: 跑步AI助手-5.0", text)
        self.assertIn("version: 0.6.0", text)
        self.assertIn("workflow:", text)


def add_test(name, fn):
    setattr(WorkflowContractTests, name, fn)


def make_node_shell_test(node):
    def test(self):
        self.assertTrue(node["id"])
        self.assertTrue(node["data"]["title"])
        self.assertTrue(node["data"]["type"])
        self.assertIsInstance(node.get("position"), dict)
        self.assertIn("x", node["position"])
        self.assertIn("y", node["position"])
    return test


def make_node_required_params_test(node):
    def test(self):
        data = node["data"]
        node_type = data["type"]
        if node_type == "start":
            self.assertIn("variables", data)
        elif node_type == "llm":
            self.assertTrue(data["model"]["name"])
            self.assertTrue(data["model"]["provider"])
            self.assertIsNotNone(data["model"]["completion_params"]["temperature"])
            self.assertTrue(data["prompt_template"])
        elif node_type == "code":
            self.assertEqual(data["code_language"], "python3")
            self.assertTrue(data["code"].strip())
            self.assertTrue(data["outputs"])
        elif node_type == "if-else":
            self.assertTrue(data["cases"])
            self.assertTrue(all(case.get("conditions") for case in data["cases"]))
        elif node_type == "http-request":
            self.assertIn(data["method"], {"get", "post"})
            self.assertTrue(data["url"])
            self.assertGreaterEqual(data["timeout"]["max_read_timeout"], 1)
            self.assertGreaterEqual(data["retry_config"]["max_retries"], 0)
        elif node_type == "knowledge-retrieval":
            self.assertTrue(data["dataset_ids"])
            self.assertEqual(data["retrieval_mode"], "multiple")
            self.assertTrue(data["query_variable_selector"])
        elif node_type == "answer":
            self.assertTrue(data["answer"])
        elif node_type == "assigner":
            self.assertTrue(data["items"])
        elif node_type == "variable-aggregator":
            self.assertTrue(data["variables"])
            self.assertEqual(data["output_type"], "string")
        else:
            self.fail(f"uncovered node type: {node_type}")
    return test


def make_edge_test(edge):
    def test(self):
        self.assertIn(edge["source"], NODE_BY_ID)
        self.assertIn(edge["target"], NODE_BY_ID)
        self.assertNotEqual(edge["source"], edge["target"])
        self.assertIn(edge.get("type"), {"custom", "default", None})
    return test


def make_if_case_test(node, case):
    def test(self):
        self.assertTrue(case["id"])
        self.assertIn(case["logical_operator"], {"and", "or"})
        for condition in case["conditions"]:
            self.assertTrue(condition["comparison_operator"])
            self.assertTrue(condition["variable_selector"])
            self.assert_selector(condition["variable_selector"])
    return test


def make_code_input_test(node, var):
    def test(self):
        self.assertTrue(var["variable"])
        self.assertTrue(var["value_type"])
        self.assert_selector(var["value_selector"])
    return test


def make_code_output_test(node, output_name, spec):
    def test(self):
        self.assertTrue(output_name)
        self.assertTrue(spec["type"])
    return test


def make_llm_prompt_test(node, prompt):
    def test(self):
        self.assertIn(prompt["role"], {"system", "user", "assistant"})
        self.assertTrue(prompt["text"].strip())
    return test


def make_answer_ref_test(node, ref):
    def test(self):
        self.assert_selector(ref.split("."))
    return test


def make_http_test(node):
    def test(self):
        data = node["data"]
        self.assertIn(data["error_strategy"], {"fail-branch", "default-value"})
        self.assertIs(data["retry_config"]["retry_enabled"], True)
        self.assertIn(data["retry_config"]["max_retries"], {2, 3})
        for ref in selectors_in_text(data["url"]):
            self.assert_selector(ref.split("."))
    return test


def make_knowledge_test(node):
    def test(self):
        data = node["data"]
        self.assertTrue(all(len(dataset_id) == 36 for dataset_id in data["dataset_ids"]))
        self.assert_selector(data["query_variable_selector"])
        self.assertTrue(data["multiple_retrieval_config"]["reranking_model"]["model"])
    return test


def make_variable_test(item):
    def test(self):
        self.assertTrue(item["name"])
        self.assertTrue(item["selector"])
        self.assertTrue(item["value_type"])
    return test


def make_assigner_test(node, item):
    def test(self):
        self.assertEqual(item["operation"], "over-write")
        self.assertIn(tuple(item["variable_selector"]), KNOWN_SELECTORS)
        self.assert_selector(item["value"])
    return test


def make_aggregator_test(node, selector):
    def test(self):
        self.assert_selector(selector)
    return test


for i, node in enumerate(NODES):
    label = safe_id(f"{node['id']}_{node['data']['title']}")
    add_test(f"test_010_node_shell_{i:03d}_{label}", make_node_shell_test(node))
    add_test(f"test_020_node_required_params_{i:03d}_{label}", make_node_required_params_test(node))

for i, edge in enumerate(EDGES):
    add_test(f"test_030_edge_{i:03d}_{safe_id(edge['id'])}", make_edge_test(edge))

case_index = 0
for node in NODES:
    if node["data"]["type"] == "if-else":
        for case in node["data"]["cases"]:
            add_test(f"test_040_if_case_{case_index:03d}_{safe_id(case['id'])}", make_if_case_test(node, case))
            case_index += 1

input_index = 0
output_index = 0
for node in NODES:
    if node["data"]["type"] == "code":
        for var in node["data"].get("variables", []):
            add_test(f"test_050_code_input_{input_index:03d}_{safe_id(node['id'] + '_' + var['variable'])}", make_code_input_test(node, var))
            input_index += 1
        for output_name, spec in node["data"].get("outputs", {}).items():
            add_test(f"test_060_code_output_{output_index:03d}_{safe_id(node['id'] + '_' + output_name)}", make_code_output_test(node, output_name, spec))
            output_index += 1

prompt_index = 0
for node in NODES:
    if node["data"]["type"] == "llm":
        for prompt in node["data"].get("prompt_template", []):
            add_test(f"test_070_llm_prompt_{prompt_index:03d}_{safe_id(node['id'] + '_' + prompt['role'])}", make_llm_prompt_test(node, prompt))
            prompt_index += 1

ref_index = 0
for node in NODES:
    if node["data"]["type"] == "answer":
        for ref in selectors_in_text(node["data"]["answer"]):
            add_test(f"test_080_answer_ref_{ref_index:03d}_{safe_id(node['id'] + '_' + ref)}", make_answer_ref_test(node, ref))
            ref_index += 1

for i, node in enumerate([node for node in NODES if node["data"]["type"] == "http-request"]):
    add_test(f"test_090_http_node_{i:03d}_{safe_id(node['data']['title'])}", make_http_test(node))

for i, node in enumerate([node for node in NODES if node["data"]["type"] == "knowledge-retrieval"]):
    add_test(f"test_100_knowledge_node_{i:03d}_{safe_id(node['data']['title'])}", make_knowledge_test(node))

for i, item in enumerate(SYSTEM_VARIABLES + CONVERSATION_VARIABLES + NODE_VARIABLES):
    add_test(f"test_110_variable_{i:03d}_{safe_id('.'.join(item['selector']))}", make_variable_test(item))

assigner_index = 0
for node in NODES:
    if node["data"]["type"] == "assigner":
        for item in node["data"].get("items", []):
            add_test(f"test_120_assigner_{assigner_index:03d}_{safe_id('.'.join(item['variable_selector']))}", make_assigner_test(node, item))
            assigner_index += 1

aggregator_index = 0
for node in NODES:
    if node["data"]["type"] == "variable-aggregator":
        for selector in node["data"].get("variables", []):
            add_test(f"test_130_aggregator_{aggregator_index:03d}_{safe_id('.'.join(selector))}", make_aggregator_test(node, selector))
            aggregator_index += 1


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