from __future__ import annotations

import json
import re
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}


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


def assert_selector(selector):
    selector = tuple(selector)
    assert selector in KNOWN_SELECTORS or selector[0] in NODE_BY_ID
    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())
        assert selector[1] in known_names, f"unknown node output selector: {selector}"


def test_app_identity():
    assert APP["id"] == "897352d5-faea-4467-9392-8cf70323c764"
    assert APP["name"] == "跑步AI助手-5.0"
    assert APP["mode"] == "advanced-chat"


def test_workflow_shape():
    assert WORKFLOW["id"] == "5f5637de-d434-4f3b-a2b1-b569afea3c89"
    assert len(NODES) == 58
    assert len(EDGES) == 72


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


import pytest


@pytest.mark.parametrize("node", NODES, ids=lambda n: f"{n['id']}:{n['data'].get('title')}")
def test_every_node_has_required_shell(node):
    assert node["id"]
    assert node["data"]["title"]
    assert node["data"]["type"]
    assert isinstance(node.get("position"), dict)
    assert "x" in node["position"] and "y" in node["position"]


@pytest.mark.parametrize("node", NODES, ids=lambda n: f"{n['id']}:{n['data'].get('type')}")
def test_every_node_type_has_required_parameters(node):
    data = node["data"]
    node_type = data["type"]
    if node_type == "start":
        assert "variables" in data
    elif node_type == "llm":
        assert data["model"]["name"]
        assert data["model"]["provider"]
        assert data["model"]["completion_params"]["temperature"] is not None
        assert data["prompt_template"]
    elif node_type == "code":
        assert data["code_language"] == "python3"
        assert data["code"].strip()
        assert data["outputs"]
    elif node_type == "if-else":
        assert data["cases"]
        assert all(case.get("conditions") for case in data["cases"])
    elif node_type == "http-request":
        assert data["method"] in {"get", "post"}
        assert data["url"]
        assert data["timeout"]["read"] >= 1
        assert data["retry_config"]["max_retries"] >= 0
    elif node_type == "knowledge-retrieval":
        assert data["dataset_ids"]
        assert data["retrieval_mode"] == "multiple"
        assert data["query_variable_selector"]
    elif node_type == "answer":
        assert data["answer"]
    elif node_type == "assigner":
        assert data["items"]
    elif node_type == "variable-aggregator":
        assert data["variables"]
        assert data["output_type"] == "string"
    else:
        raise AssertionError(f"uncovered node type: {node_type}")


@pytest.mark.parametrize("edge", EDGES, ids=lambda e: e["id"])
def test_every_edge_connects_existing_nodes(edge):
    assert edge["source"] in NODE_BY_ID
    assert edge["target"] in NODE_BY_ID
    assert edge["source"] != edge["target"]
    assert edge.get("type") in {"custom", "default", None}


IF_CASES = [
    (node, case)
    for node in NODES
    if node["data"]["type"] == "if-else"
    for case in node["data"]["cases"]
]


@pytest.mark.parametrize("node,case", IF_CASES, ids=lambda x: x[1]["id"])
def test_every_if_else_case_has_valid_conditions(node, case):
    assert case["id"]
    assert case["logical_operator"] in {"and", "or"}
    for condition in case["conditions"]:
        assert condition["comparison_operator"]
        assert condition["variable_selector"]
        assert_selector(condition["variable_selector"])


CODE_INPUTS = [
    (node, var)
    for node in NODES
    if node["data"]["type"] == "code"
    for var in node["data"].get("variables", [])
]


@pytest.mark.parametrize("node,var", CODE_INPUTS, ids=lambda x: f"{x[0]['id']}:{x[1]['variable']}")
def test_every_code_input_selector_is_valid(node, var):
    assert var["variable"]
    assert var["value_type"]
    assert_selector(var["value_selector"])


CODE_OUTPUTS = [
    (node, output_name, spec)
    for node in NODES
    if node["data"]["type"] == "code"
    for output_name, spec in node["data"].get("outputs", {}).items()
]


@pytest.mark.parametrize("node,output_name,spec", CODE_OUTPUTS, ids=lambda x: f"{x[0]['id']}:{x[1]}")
def test_every_code_output_is_typed(node, output_name, spec):
    assert output_name
    assert spec["type"]


LLM_PROMPTS = [
    (node, prompt)
    for node in NODES
    if node["data"]["type"] == "llm"
    for prompt in node["data"].get("prompt_template", [])
]


@pytest.mark.parametrize("node,prompt", LLM_PROMPTS, ids=lambda x: f"{x[0]['id']}:{x[1]['role']}")
def test_every_llm_prompt_is_configured(node, prompt):
    assert prompt["role"] in {"system", "user", "assistant"}
    assert prompt["text"].strip()


ANSWER_REFS = [
    (node, ref)
    for node in NODES
    if node["data"]["type"] == "answer"
    for ref in selectors_in_text(node["data"]["answer"])
]


@pytest.mark.parametrize("node,ref", ANSWER_REFS, ids=lambda x: f"{x[0]['id']}:{x[1]}")
def test_answer_references_are_known(node, ref):
    assert_selector(ref.split("."))


HTTP_NODES = [node for node in NODES if node["data"]["type"] == "http-request"]


@pytest.mark.parametrize("node", HTTP_NODES, ids=lambda n: n["data"]["title"])
def test_http_nodes_have_safe_retry_and_error_strategy(node):
    data = node["data"]
    assert data["error_strategy"] in {"fail-branch", "default-value"}
    assert data["retry_config"]["retry_enabled"] is True
    assert data["retry_config"]["max_retries"] in {2, 3}
    for ref in selectors_in_text(data["url"]):
        assert_selector(ref.split("."))


KNOWLEDGE_NODES = [node for node in NODES if node["data"]["type"] == "knowledge-retrieval"]


@pytest.mark.parametrize("node", KNOWLEDGE_NODES, ids=lambda n: n["data"]["title"])
def test_knowledge_nodes_have_dataset_and_query(node):
    data = node["data"]
    assert all(len(dataset_id) == 36 for dataset_id in data["dataset_ids"])
    assert_selector(data["query_variable_selector"])
    assert data["multiple_retrieval_config"]["reranking_model"]["model"]


ALL_VARIABLES = SYSTEM_VARIABLES + CONVERSATION_VARIABLES + NODE_VARIABLES


@pytest.mark.parametrize("item", ALL_VARIABLES, ids=lambda i: ".".join(i["selector"]))
def test_every_registered_variable_has_selector_and_type(item):
    assert item["name"]
    assert item["selector"]
    assert item["value_type"]


ASSIGNER_ITEMS = [
    (node, item)
    for node in NODES
    if node["data"]["type"] == "assigner"
    for item in node["data"].get("items", [])
]


@pytest.mark.parametrize("node,item", ASSIGNER_ITEMS, ids=lambda x: ".".join(x[1]["variable_selector"]))
def test_assigner_items_write_known_conversation_variables(node, item):
    assert item["operation"] == "over-write"
    assert tuple(item["variable_selector"]) in KNOWN_SELECTORS
    assert_selector(item["value"])


AGGREGATOR_INPUTS = [
    (node, selector)
    for node in NODES
    if node["data"]["type"] == "variable-aggregator"
    for selector in node["data"].get("variables", [])
]


@pytest.mark.parametrize("node,selector", AGGREGATOR_INPUTS, ids=lambda x: ".".join(x[1]))
def test_aggregator_inputs_are_known(node, selector):
    assert_selector(selector)
