import pytest
from app.modules.pipeline._dag_levels import compute_dag_levels


def test_all_roots_single_level():
    items = {"a": {"parent_id": ""}, "b": {"parent_id": ""}, "c": {"parent_id": ""}}
    order = ["a", "b", "c"]
    renderable = {"a", "b", "c"}
    levels = compute_dag_levels(order, items, renderable)
    assert levels == [["a", "b", "c"]]


def test_linear_chain_4_levels():
    items = {
        "n0": {"parent_id": ""},
        "n1": {"parent_id": "n0"},
        "n2": {"parent_id": "n1"},
        "n3": {"parent_id": "n2"},
    }
    order = ["n0", "n1", "n2", "n3"]
    renderable = set(order)
    levels = compute_dag_levels(order, items, renderable)
    assert levels == [["n0"], ["n1"], ["n2"], ["n3"]]


def test_fanout_then_merge():
    items = {
        "root": {"parent_id": ""},
        "a": {"parent_id": "root"},
        "b": {"parent_id": "root"},
        "merge": {"parent_id": "a"},
    }
    order = ["root", "a", "b", "merge"]
    renderable = set(order)
    levels = compute_dag_levels(order, items, renderable)
    assert levels == [["root"], ["a", "b"], ["merge"]]


def test_parent_outside_renderable_treated_as_root():
    items = {"a": {"parent_id": "skipped_node"}, "b": {"parent_id": ""}}
    order = ["a", "b"]
    renderable = {"a", "b"}  # "skipped_node" not renderable
    levels = compute_dag_levels(order, items, renderable)
    assert levels == [["a", "b"]]


def test_empty_input():
    assert compute_dag_levels([], {}, set()) == []


def test_cycle_flushed_as_final_batch(caplog):
    import logging
    items = {"a": {"parent_id": "b"}, "b": {"parent_id": "a"}}
    order = ["a", "b"]
    renderable = {"a", "b"}
    with caplog.at_level(logging.WARNING):
        levels = compute_dag_levels(order, items, renderable)
    assert levels == [["a", "b"]]
    assert any("unresolved" in r.message.lower() for r in caplog.records)


def test_preserves_order_within_level():
    items = {"z": {"parent_id": ""}, "a": {"parent_id": ""}, "m": {"parent_id": ""}}
    order = ["z", "a", "m"]
    renderable = set(order)
    levels = compute_dag_levels(order, items, renderable)
    assert levels == [["z", "a", "m"]]
