"""The per-parent cap on line records: capped tables keep their whole-table
chunk, and everything that is not a line record passes through."""

from quber.playground.retrieval import RetrievedChunk, _cap_line_records


def chunk(chunk_id, chunk_type, parent=None):
    return RetrievedChunk(
        chunk_id=chunk_id, chunk_type=chunk_type, page=0, content="", score=0.0, parent_chunk_id=parent
    )


def test_cap_keeps_two_line_records_per_parent_in_order():
    pool = [
        chunk("row1", "line_item", parent="tableA"),
        chunk("row2", "line_item", parent="tableA"),
        chunk("row3", "line_item", parent="tableA"),
        chunk("row4", "line_item", parent="tableB"),
    ]
    kept = _cap_line_records(pool, cap=2)
    assert [c.chunk_id for c in kept] == ["row1", "row2", "row4"]


def test_whole_table_chunk_is_never_capped():
    pool = [
        chunk("row1", "line_item", parent="tableA"),
        chunk("row2", "line_item", parent="tableA"),
        chunk("row3", "line_item", parent="tableA"),
        chunk("tableA", "table"),
    ]
    kept = _cap_line_records(pool, cap=2)
    assert "tableA" in [c.chunk_id for c in kept]


def test_non_line_records_pass_through():
    pool = [chunk("t", "text"), chunk("f", "figure"), chunk("orphan", "line_item", parent=None)]
    assert _cap_line_records(pool, cap=1) == pool
