Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 35 additions & 1 deletion pageindex/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1257,7 +1257,7 @@ def create_node_mapping(tree, include_page_ranges=False, max_page=None):
"end_index"} (end = next node's page_index, or max_page for the last node)."""
def get_all_nodes(tree):
if isinstance(tree, dict):
return [tree] + [node for child in tree.get('nodes', []) for node in get_all_nodes(child)]
return [tree] + [node for child in tree.get('nodes') or [] for node in get_all_nodes(child)]
elif isinstance(tree, list):
return [node for item in tree for node in get_all_nodes(item)]
return []
Expand All @@ -1276,6 +1276,40 @@ def get_all_nodes(tree):
}
return mapping

def _require_node_tree(tree):
if not (isinstance(tree, list) or (isinstance(tree, dict) and 'node_id' in tree)):
raise TypeError("tree must be a node list such as get_document_structure(doc_id) "
"or get_tree(doc_id)['result'], not the whole get_tree response; "
f"got {type(tree).__name__}")

def get_node_path(tree, node_id):
"""[top-level ancestor, ..., node] for node_id; [] if absent."""
_require_node_tree(tree)
if not isinstance(node_id, str):
raise TypeError(f"node_id must be a str like '0007', got {node_id!r}")
for node in [tree] if isinstance(tree, dict) else tree:
if node.get('node_id') == node_id:
return [node]
path = get_node_path(node.get('nodes') or [], node_id)
if path:
return [node] + path
return []

def get_node(tree, node_id):
"""The node with node_id, or None."""
path = get_node_path(tree, node_id)
return path[-1] if path else None

def get_node_parent(tree, node_id):
"""The parent of node_id; None for a top-level or absent node."""
path = get_node_path(tree, node_id)
return path[-2] if len(path) > 1 else None

def get_node_map(tree):
"""{node_id: node} for every node in tree."""
_require_node_tree(tree)
return create_node_mapping(tree)

def print_tree(tree, exclude_fields=None, indent=0):
"""Outline view; passing exclude_fields gives the 0.2.8 pprint view."""
if exclude_fields is not None:
Expand Down
35 changes: 34 additions & 1 deletion tests/test_package_surface.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,11 @@
import subprocess
import sys

from pageindex.utils import create_node_mapping, print_tree, remove_fields
import pytest

from pageindex.utils import (create_node_mapping, get_node, get_node_map,
get_node_parent, get_node_path, print_tree,
remove_fields)

TREE = [
{"title": "Root", "node_id": "0000", "page_index": 1,
Expand Down Expand Up @@ -47,6 +51,35 @@ def test_print_tree_exclude_fields(capsys):
assert "[0000] Root" in capsys.readouterr().out


# ── tree navigation ──

def test_node_navigation():
root, child, tail = TREE[0], TREE[0]["nodes"][0], TREE[1]
assert get_node(TREE, "0001") is child
assert get_node(TREE, "9999") is None
assert get_node_parent(TREE, "0001") is root
assert get_node_parent(TREE, "0002") is None
assert get_node_path(TREE, "0001") == [root, child]
assert get_node_path(TREE, "0002") == [tail]
assert get_node_path(TREE, "9999") == []
assert get_node(root, "0001") is child
assert get_node_map(TREE) == {"0000": root, "0001": child, "0002": tail}
leaf = {"node_id": "0000", "nodes": None}
assert get_node_map([leaf]) == {"0000": leaf}


def test_node_navigation_rejects_wrong_input():
envelope = {"doc_id": "d", "status": "completed", "result": TREE}
with pytest.raises(TypeError, match="got dict"):
get_node(envelope, "0001")
with pytest.raises(TypeError, match="got dict"):
get_node_map(envelope)
with pytest.raises(TypeError, match="got NoneType"):
get_node_map(None)
with pytest.raises(TypeError, match="node_id"):
get_node(TREE, 1)


# ── import cost: the SDK must not pay for the indexing stack ──

def test_import_pageindex_is_lazy():
Expand Down
Loading