diff --git a/project/paperbench/paperbench/rubric/tasks.py b/project/paperbench/paperbench/rubric/tasks.py index 887e6984..a675092e 100644 --- a/project/paperbench/paperbench/rubric/tasks.py +++ b/project/paperbench/paperbench/rubric/tasks.py @@ -270,7 +270,7 @@ def get_prior_nodes(self, root: TaskNode, max_prior_nodes: int | None = None) -> required_nodes = required_nodes[:-1] # Don't include the target node if max_prior_nodes is not None: - required_nodes = required_nodes[-max_prior_nodes:] + required_nodes = required_nodes[-max_prior_nodes:] if max_prior_nodes > 0 else [] return required_nodes def get_descendants_depth_first(self) -> list[Self]: diff --git a/project/paperbench/tests/unit/test_rubric_prior_nodes.py b/project/paperbench/tests/unit/test_rubric_prior_nodes.py new file mode 100644 index 00000000..056b5a71 --- /dev/null +++ b/project/paperbench/tests/unit/test_rubric_prior_nodes.py @@ -0,0 +1,33 @@ +from paperbench.rubric.tasks import TaskNode + + +def _rubric() -> tuple[TaskNode, TaskNode]: + b = TaskNode( + id="B", + requirements="B", + weight=1, + task_category="Code Development", + ) + f = TaskNode( + id="F", + requirements="F", + weight=1, + task_category="Code Development", + ) + g = TaskNode( + id="G", + requirements="G", + weight=1, + task_category="Code Development", + ) + c = TaskNode(id="C", requirements="C", weight=1, sub_tasks=[f, g]) + root = TaskNode(id="A", requirements="A", weight=1, sub_tasks=[b, c]) + return root, g + + +def test_get_prior_nodes_respects_zero_limit() -> None: + root, target = _rubric() + + assert [node.id for node in target.get_prior_nodes(root)] == ["A", "B", "C", "F"] + assert target.get_prior_nodes(root, max_prior_nodes=0) == [] + assert [node.id for node in target.get_prior_nodes(root, max_prior_nodes=2)] == ["C", "F"]