diff --git a/pageindex/client.py b/pageindex/client.py index 0d77269bf..c2cdd68d8 100644 --- a/pageindex/client.py +++ b/pageindex/client.py @@ -152,7 +152,8 @@ def _needs_model(surface: str) -> PageIndexAPIError: _LOCAL_INDEX_KEYS = ("model", "summary_model", "backend", "storage_path", "summary_max_words", "summary_concurrency", - "use_embedded_toc", "optimize") + "summary_max_input_tokens", "summary_scope", "use_embedded_toc", + "optimize") # Near-synonyms of "cloud" that would otherwise parse as model names — # a silent wrong mode. They error, pointing at the real word. @@ -178,7 +179,9 @@ def _env_cloud_key(spelling: str, inline: str = "api_key=...") -> str: _ARG_TYPES: "dict[str, tuple[type, ...]]" = { "model": (str,), "index_model": (str,), "summary_model": (str,), "chat_model": (str,), "retrieve_model": (str,), "summary_max_words": (int,), - "summary_concurrency": (int,), "use_embedded_toc": (bool,), "optimize": (str,), + "summary_concurrency": (int,), "summary_max_input_tokens": (int,), + "summary_scope": (str,), + "use_embedded_toc": (bool,), "optimize": (str,), "storage_path": (str, os.PathLike), "index_backend": (dict,), "chat_backend": (dict,)} @@ -356,7 +359,8 @@ class PageIndexClient: environment), ``"local"``, a local index model name, or a dict: ``{"api_key": ...}`` for cloud, ``{"model", "summary_model", "backend", "storage_path", - "summary_max_words", "summary_concurrency", "use_embedded_toc", + "summary_max_words", "summary_concurrency", + "summary_max_input_tokens", "summary_scope", "use_embedded_toc", "optimize"}`` for local. An optional ``"mode"`` key (``"cloud"`` / ``"local"``) states the side and must agree with the other keys; ``{"mode": @@ -411,6 +415,18 @@ class PageIndexClient: expand up to its own ceiling of 32. The lanes overlap, so up to cap + min(32, cap) calls run at once. Defaults to 64. A ``mode="standard"`` submit refuses either summary knob. + summary_max_input_tokens (int, optional): Local mode only - the + context size of the indexing model. Prompts that grow with the + document stay within it: the document description is cut from + its deepest level, a flash leaf too long for one call is + summarized in parts, and an expand prompt that would overrun it + is skipped. Defaults to unbounded. + summary_scope (str, optional): Local flash mode only - ``"pages"`` + summarizes a leaf from the pages of its node, ``"section"`` from + the layout blocks between its heading and the next located one, + which leaves out the end of the previous section and the start + of the next. Defaults to ``"pages"``; a ``mode="standard"`` + submit refuses ``"section"``. use_embedded_toc (bool, optional): Local mode only — whether flash indexing consumes the PDF's embedded bookmarks when they look trustworthy. Defaults to True. @@ -467,6 +483,8 @@ def __init__( summary_model: Optional[str] = None, summary_max_words: Optional[int] = None, summary_concurrency: Optional[int] = None, + summary_max_input_tokens: Optional[int] = None, + summary_scope: Optional[str] = None, use_embedded_toc: Optional[bool] = None, optimize: Optional[str] = None, retrieve_model: Optional[str] = None, @@ -496,6 +514,8 @@ def __init__( ("summary_model", summary_model), ("summary_max_words", summary_max_words), ("summary_concurrency", summary_concurrency), + ("summary_max_input_tokens", summary_max_input_tokens), + ("summary_scope", summary_scope), ("use_embedded_toc", use_embedded_toc), ("optimize", optimize), ("index_backend", index_backend), @@ -582,10 +602,14 @@ def __init__( raise PageIndexAPIError( f"{shown} is empty — it configures nothing. Pass a " "real value, or drop the argument.") + if name == "summary_scope" and value not in ("pages", "section"): + raise PageIndexAPIError( + f'{shown} must be "pages" or "section", got {value!r}.') if name == "optimize" and value not in ("full", "merge", "off"): raise PageIndexAPIError( f'{shown} must be "full", "merge" or "off", got {value!r}.') - if (name in ("summary_max_words", "summary_concurrency") + if (name in ("summary_max_words", "summary_concurrency", + "summary_max_input_tokens") and isinstance(value, int) and value < 1): raise PageIndexAPIError( f"{shown} must be a positive int, got {value!r}.") @@ -660,6 +684,8 @@ def __init__( index_backend=index_conf.get("index_backend"), summary_max_words=index_conf.get("summary_max_words"), summary_concurrency=index_conf.get("summary_concurrency"), + summary_max_input_tokens=index_conf.get("summary_max_input_tokens"), + summary_scope=index_conf.get("summary_scope", "pages"), use_embedded_toc=index_conf.get("use_embedded_toc", True), optimize=index_conf.get("optimize", "full"), ) @@ -2698,6 +2724,8 @@ def __init__( summary_model: Optional[str] = None, summary_max_words: Optional[int] = None, summary_concurrency: Optional[int] = None, + summary_max_input_tokens: Optional[int] = None, + summary_scope: Optional[str] = None, use_embedded_toc: Optional[bool] = None, optimize: Optional[str] = None, retrieve_model: Optional[str] = None, @@ -2711,6 +2739,8 @@ def __init__( model=model, summary_model=summary_model, summary_max_words=summary_max_words, summary_concurrency=summary_concurrency, + summary_max_input_tokens=summary_max_input_tokens, + summary_scope=summary_scope, use_embedded_toc=use_embedded_toc, optimize=optimize, retrieve_model=retrieve_model, storage_path=storage_path, index_backend=index_backend, chat_backend=chat_backend, diff --git a/pageindex/flash/api.py b/pageindex/flash/api.py index 5e3e397e1..34d693e07 100644 --- a/pageindex/flash/api.py +++ b/pageindex/flash/api.py @@ -73,14 +73,16 @@ def _validate_pdf(pdf): return pdf -async def _summarize(structure, page_list, model, concurrency=None, max_words=None): +async def _summarize(structure, page_list, model, concurrency=None, max_words=None, + max_input_tokens=None, blocks=None): from ..utils import summarize_tree await summarize_tree(structure, page_list, model=model, concurrency=concurrency, - max_words=max_words) + max_words=max_words, max_input_tokens=max_input_tokens, + blocks=blocks) async def _optimize_async(structure, page_texts, do_expand, model, on_final=None, - concurrency=None): + concurrency=None, max_input_tokens=None): """Merge/expand refinement after extraction, overlapped with the summaries when `on_final` is passed; without it the caller runs them after. @@ -92,7 +94,8 @@ async def _optimize_async(structure, page_texts, do_expand, model, on_final=None lines = _page_lines(page_texts) outcome = await optimize(structure, page_texts, lines, model=model, do_expand=do_expand, page_count=len(page_texts), - on_final=on_final, concurrency=concurrency) + on_final=on_final, concurrency=concurrency, + max_input_tokens=max_input_tokens) return {"merges": outcome["merges"], "expands": outcome["expands"], "same_page_merges": outcome["same_page_merges"], "same_page_dropped": outcome["same_page_dropped"], @@ -107,16 +110,19 @@ def _optimize(structure, page_texts, do_expand, model, concurrency=None): async def _optimize_and_summarize(structure, page_texts, optimize_model, summary_model, - concurrency, max_words=None): + concurrency, max_words=None, max_input_tokens=None, + blocks=None): """Expand and summarize on one loop: a node is summarized as soon as expand can no longer change it, a parent once its children are done.""" from ..utils import SummaryScheduler scheduler = SummaryScheduler(structure, [(text, 0) for text in page_texts], model=summary_model, concurrency=concurrency, - max_words=max_words) + max_words=max_words, max_input_tokens=max_input_tokens, + blocks=blocks) report = await _optimize_async(structure, page_texts, True, optimize_model, on_final=scheduler.mark_final, - concurrency=concurrency) + concurrency=concurrency, + max_input_tokens=max_input_tokens) await scheduler.finish() return report @@ -180,12 +186,16 @@ def flash_rejection_reason(result: dict, standard_hint: str = "mode='standard'") def page_index_flash(pdf, summary=True, summary_model=None, optimize: str | bool | None = None, optimize_expand=None, optimize_model=None, summary_concurrency=None, - use_embedded_toc=True, summary_max_words=None) -> dict: - """Build a PageIndex tree structure from a PDF using layout statistics. The tree extraction itself uses no LLM; by default an LLM writes node summaries and expands the tree (``summary=False, optimize=False`` runs fully LLM-free). Args: pdf: path to a PDF file (``str`` or ``pathlib.Path``) or an in-memory binary stream (``io.BytesIO``). summary: if True, generate LLM summaries for each node (requires ``summary_model``). summary_model: the LLM model identifier to use for summary generation. optimize: ``"full"`` for merge + LLM expand (a model unreachable after the retry ladder — a missing credential included — fails the run loudly from expand itself; a per-prompt rejection leaves just that node collapsed), ``"merge"`` for deterministic merge only, ``False`` to disable. ``True`` is accepted as ``"full"`` for backward compatibility; defaults to ``"full"``. Expand needs readable page text, so a bookmark-only or scanned PDF runs the merge half only (``expands`` reports 0). optimize_expand: deprecated — use ``optimize``. Honored only when ``optimize`` is not passed (or is the legacy ``True``): ``False`` maps to ``"merge"``, ``True`` to ``"full"``. optimize_model: the LLM model for expand (defaults to the summary model). summary_concurrency: cap on simultaneous indexing model calls per lane: the summaries, and expand up to its own ceiling of 32 (the lanes overlap, so up to cap + min(32, cap) calls run at once); None uses the library defaults (64 and 32). use_embedded_toc: if True, consume the PDF's embedded bookmarks when trustworthy: deep bookmarks become the frame and the detected sections they lack are grafted back in after noise filtering, coarse ones become the chapter frame with detected nodes re-hung under them (deeper sparse entries are filled in when the page text confirms them, and garbled extracted titles are repaired from the bookmark strings), garbage ones are ignored. On by default; pass False for the pure detected structure. summary_max_words: word cap each model-written node summary is asked to stay within (short leaves keep their raw text); None uses the library default (150). Returns: dict with keys ``doc_name``, ``doc_title``, ``structure`` (a list of ``{"title", "node_id", "start_index", "end_index"}`` dicts; ``"nodes"`` holds the children where there are any and ``"summary"`` appears when summaries ran; page indexes are 1-based; a hierarchy that starts after page 1 is preceded by a ``Preface`` node covering the pages before it, as in standard mode, and the first section's page too unless that heading opens it; a parent whose first child starts on a later page opens with a child titled ``" (intro)"`` holding the pages before it; a parent's range and summary cover its whole subtree) and ``has_abstract_or_references_section`` (True when a top-level entry is an abstract or references heading). ``toc_source`` says where the structure came from: ``"detected"`` (layout), ``"bookmarks"`` (the embedded outline), ``"hybrid"`` (bookmarks framing the detected sections), ``"pages"`` (no hierarchy found, so one node per page titled ``Page N``; left unsummarized and unoptimized when there are more than ``FLAT_TREE_MAX_NODES`` pages, a size the local client and CLI refuse) or ``"unreadable"`` (no page carries text; ``structure`` is empty). With ``optimize`` an ``optimize`` key reports merge/expand counts and before/after search-cost metrics; a refused flat tree carries neither it nor node summaries. """ + use_embedded_toc=True, summary_max_words=None, + summary_max_input_tokens=None, summary_scope="pages") -> dict: + """Build a PageIndex tree structure from a PDF using layout statistics. The tree extraction itself uses no LLM; by default an LLM writes node summaries and expands the tree (``summary=False, optimize=False`` runs fully LLM-free). Args: pdf: path to a PDF file (``str`` or ``pathlib.Path``) or an in-memory binary stream (``io.BytesIO``). summary: if True, generate LLM summaries for each node (requires ``summary_model``). summary_model: the LLM model identifier to use for summary generation. optimize: ``"full"`` for merge + LLM expand (a model unreachable after the retry ladder — a missing credential included — fails the run loudly from expand itself; a per-prompt rejection leaves just that node collapsed), ``"merge"`` for deterministic merge only, ``False`` to disable. ``True`` is accepted as ``"full"`` for backward compatibility; defaults to ``"full"``. Expand needs readable page text, so a bookmark-only or scanned PDF runs the merge half only (``expands`` reports 0). optimize_expand: deprecated — use ``optimize``. Honored only when ``optimize`` is not passed (or is the legacy ``True``): ``False`` maps to ``"merge"``, ``True`` to ``"full"``. optimize_model: the LLM model for expand (defaults to the summary model). summary_concurrency: cap on simultaneous indexing model calls per lane: the summaries, and expand up to its own ceiling of 32 (the lanes overlap, so up to cap + min(32, cap) calls run at once); None uses the library defaults (64 and 32). use_embedded_toc: if True, consume the PDF's embedded bookmarks when trustworthy: deep bookmarks become the frame and the detected sections they lack are grafted back in after noise filtering, coarse ones become the chapter frame with detected nodes re-hung under them (deeper sparse entries are filled in when the page text confirms them, and garbled extracted titles are repaired from the bookmark strings), garbage ones are ignored. On by default; pass False for the pure detected structure. summary_max_words: word cap each model-written node summary is asked to stay within (short leaves keep their raw text); None uses the library default (150). summary_max_input_tokens: context size of the indexing model; the summary and expand prompts are kept within it (a leaf too long for one call is summarized in parts, an expand that overruns it is skipped); None leaves them unbounded. summary_scope: ``"pages"`` summarizes a leaf from the pages of its node, ``"section"`` from the layout blocks between its heading and the next located one (a leaf falls back to its pages when a section without a located heading may start inside it). Returns: dict with keys ``doc_name``, ``doc_title``, ``structure`` (a list of ``{"title", "node_id", "start_index", "end_index"}`` dicts; ``"nodes"`` holds the children where there are any and ``"summary"`` appears when summaries ran; page indexes are 1-based; a hierarchy that starts after page 1 is preceded by a ``Preface`` node covering the pages before it, as in standard mode, and the first section's page too unless that heading opens it; a parent whose first child starts on a later page opens with a child titled ``" (intro)"`` holding the pages before it; a parent's range and summary cover its whole subtree) and ``has_abstract_or_references_section`` (True when a top-level entry is an abstract or references heading). ``toc_source`` says where the structure came from: ``"detected"`` (layout), ``"bookmarks"`` (the embedded outline), ``"hybrid"`` (bookmarks framing the detected sections), ``"pages"`` (no hierarchy found, so one node per page titled ``Page N``; left unsummarized and unoptimized when there are more than ``FLAT_TREE_MAX_NODES`` pages, a size the local client and CLI refuse) or ``"unreadable"`` (no page carries text; ``structure`` is empty). With ``optimize`` an ``optimize`` key reports merge/expand counts and before/after search-cost metrics; a refused flat tree carries neither it nor node summaries. """ for name, value in (("summary_concurrency", summary_concurrency), - ("summary_max_words", summary_max_words)): + ("summary_max_words", summary_max_words), + ("summary_max_input_tokens", summary_max_input_tokens)): if value is not None and not (isinstance(value, numbers.Integral) and int(value) >= 1): raise ValueError(f"{name} must be a positive int, got {value!r}") + if summary_scope not in ("pages", "section"): + raise ValueError(f'summary_scope must be "pages" or "section", got {summary_scope!r}') if optimize_expand is not None: import warnings warnings.warn( @@ -202,7 +212,9 @@ def page_index_flash(pdf, summary=True, summary_model=None, elif optimize not in ("full", "merge"): raise ValueError( f"optimize must be 'full', 'merge', or False, got {optimize!r}") - result = extract_toc(_validate_pdf(pdf), use_embedded_toc=use_embedded_toc) + result = extract_toc(_validate_pdf(pdf), use_embedded_toc=use_embedded_toc, + **({"with_blocks": True} if summary_scope == "section" else {})) + blocks = result.pop("block_texts", None) structure = result.get("structure", []) if not structure: # the layout yields no hierarchy; the pages themselves are the tree @@ -230,7 +242,8 @@ def page_index_flash(pdf, summary=True, summary_model=None, result["optimize"] = asyncio.run(_optimize_and_summarize( structure, pages, optimize_model=optimize_model or summary_model, summary_model=summary_model, concurrency=summary_concurrency, - max_words=summary_max_words)) + max_words=summary_max_words, max_input_tokens=summary_max_input_tokens, + blocks=blocks)) return result if optimize and structure: result["optimize"] = _optimize(structure, pages, do_expand, @@ -241,7 +254,9 @@ def page_index_flash(pdf, summary=True, summary_model=None, page_list = [(text, 0) for text in pages] asyncio.run(_summarize(structure, page_list, summary_model, concurrency=summary_concurrency, - max_words=summary_max_words)) + max_words=summary_max_words, + max_input_tokens=summary_max_input_tokens, + blocks=blocks)) elif structure: from ..utils import strip_internal_keys strip_internal_keys(structure) # summarize_tree does this on its way out diff --git a/pageindex/flash/main.py b/pageindex/flash/main.py index 3f8b3d9c4..77072d31e 100644 --- a/pageindex/flash/main.py +++ b/pageindex/flash/main.py @@ -17,7 +17,7 @@ # (re is used by the title-reject regex below) from .blocks import cluster_lines_into_blocks, BlockClusterContext -from .classification import is_body_paragraph, detect_header_footer, HeaderFooterContext, mark_watermarks, mark_toc_and_boilerplate +from .classification import is_body_paragraph, detect_header_footer, HeaderFooterContext, mark_watermarks, mark_toc_and_boilerplate, bounded_edit_distance from .labels import detect_captions, build_caption_regions, CaptionContext from .model import Rect, numbering_kind, block_text, deaccented_text, Block from .outline_assembly import ( @@ -116,6 +116,59 @@ def page_by_block_lookup(pages, block) -> Optional[PageView]: return None +# --------------------------------------------------------------------------- # +# Heading positions # +# --------------------------------------------------------------------------- # + +HEADING_TYPES = (7, 8) # unnumbered and numbered outline headings (mark_outline_block_types) + + +def _loose(text: str) -> str: + return " ".join(re.sub(r"\W+", " ", unicodedata.normalize("NFKC", text).lower()).split()) + + +def _find_heading(title: str, candidates: list) -> Optional[int]: + """Position of the block printing `title`: the closest match among heading blocks first, then + among blocks whose text is mostly the title (exact, numbered, prefixed, then near spelling).""" + wanted = _loose(title) + for headings_only in (True, False): + best = None + for pos, text, kind in candidates: + if not text or (headings_only and kind not in HEADING_TYPES) or (kind not in HEADING_TYPES and len(text) > 1.6 * len(wanted)): + continue + head = text[:len(wanted) + 14] + reach = min(len(wanted), len(head)) + rank = (0 if text == wanted else 1 if text.endswith(wanted) else 2 if text.startswith(wanted) + else 3 if bounded_edit_distance(head[:reach], wanted[:reach], 0.2 * reach) < 0.2 * reach + else None) + if rank is not None and (best is None or rank < best[0]): + best = (rank, pos) + if best is not None: + return best[1] + return None + + +def locate_headings(structure: list[dict], body: list[Block], block_pages: list[int]) -> None: + """Give each node without `_pos` the position of the block that prints its title: on its + start page, or a heading block on the page after or before it (bookmarks can be one page off).""" + by_page: dict[int, list] = {} + for pos, (block, page_no) in enumerate(zip(body, block_pages)): + by_page.setdefault(page_no, []).append((pos, _loose(block_text(block)), block.type)) + stack = list(structure) + while stack: + node = stack.pop() + stack.extend(node.get("nodes") or []) + if node.get("_pos") is not None or not _loose(node["title"]): + continue + start = node["start_index"] + found = _find_heading(node["title"], by_page.get(start, [])) + for near in (start + 1, start - 1): + if found is None: + found = _find_heading(node["title"], [c for c in by_page.get(near, []) if c[2] in HEADING_TYPES]) + if found is not None: + node["_pos"] = found + + # --------------------------------------------------------------------------- # # End-to-end entry point # # --------------------------------------------------------------------------- # @@ -125,6 +178,7 @@ def extract_toc( doc_handle: Union[str, Path, BytesIO], workers: Optional[int] = None, use_embedded_toc: bool = True, + with_blocks: bool = False, ) -> dict: """Run the full pipeline. Returns a dict shaped like:: { "doc_name": "...", "doc_title": "...", "structure": [ {"title": "...", "start_index": 1, "end_index": 3, "nodes": [...]}, ... ], "has_abstract_or_references_section": False } ``has_abstract_or_references_section`` is True when any TOP-LEVEL outline entry is an abstract-keyword heading or carries the prominent-heading flag (a references-keyword heading, plain or numbered). The valid-outline branch reports False. ``workers`` sets the process count for the per-page parallel parser: None = auto (CPU count - 1), 1 forces the sequential path; output is identical either way. ``use_embedded_toc`` consumes the PDF's embedded bookmarks when trustworthy: deep bookmarks become the frame with the detected sections they lack grafted back in, coarse ones become the chapter frame with detected nodes re-hung under them, garbage ones are ignored. On by default; pass False for the pure detected structure. ``toc_source`` is always present: ``"detected"``, ``"bookmarks"``, or ``"hybrid"``. """ # ----- 1) Parse PDF -> flat spans per page -------------------------- @@ -269,8 +323,16 @@ def extract_toc( ): outline_nodes = [] has_abstract_or_references = has_table_or_prominent(outline_nodes) + body, block_pages, block_pos = [], [], None + if with_blocks: + for page in pages: + for block in sorted(page.secondary_slot or [], key=lambda b: b.reading_order_index): + if block.type == 0 or block.type in HEADING_TYPES: + body.append(block) + block_pages.append(page.page_index) + block_pos = {id(block): index for index, block in enumerate(body)} if outline_nodes: - structure = outline_to_dict_tree(outline_nodes, total_pages=len(pages)) + structure = outline_to_dict_tree(outline_nodes, total_pages=len(pages), block_pos=block_pos) else: structure = [] @@ -300,6 +362,9 @@ def extract_toc( result["structure"], result["toc_source"] = apply_embedded_toc( structure, doc_handle, len(pages), page_texts=page_texts ) + if with_blocks: + result["block_texts"] = [block_text(block) for block in body] + locate_headings(result["structure"], body, block_pages) return result diff --git a/pageindex/flash/outline_assembly/assembly.py b/pageindex/flash/outline_assembly/assembly.py index 816c5411f..f0eed4548 100644 --- a/pageindex/flash/outline_assembly/assembly.py +++ b/pageindex/flash/outline_assembly/assembly.py @@ -244,7 +244,8 @@ def _heading_appears_at_page_top(heading: HeadingCandidate) -> bool: return True -def outline_to_dict_tree(outline_node_list: list[OutlineNode], total_pages: int) -> list[dict]: +def outline_to_dict_tree(outline_node_list: list[OutlineNode], total_pages: int, + block_pos: Optional[dict] = None) -> list[dict]: """Convert the outline tree directly to PageIndex JSON shape. Preserves the natural outline nesting without font-overlay rewriting. """ flat_nodes: list[dict] = [] @@ -276,6 +277,8 @@ def _walk_nodes(items: list[OutlineNode]) -> list[dict]: "nodes": _walk_nodes(item.child_nodes) if item.child_nodes else [], "_appear_start": _heading_appears_at_page_top(item.heading), } + if block_pos is not None: + node["_pos"] = block_pos.get(id(item.heading.group_slot)) flat_nodes.append(node) result.append(node) return result diff --git a/pageindex/local_api.py b/pageindex/local_api.py index 4802634b8..286f6bc48 100644 --- a/pageindex/local_api.py +++ b/pageindex/local_api.py @@ -41,6 +41,8 @@ def __init__(self, storage_path: str, model: str, summary_model: str, index_backend: dict | None = None, summary_max_words: int | None = None, summary_concurrency: int | None = None, + summary_max_input_tokens: int | None = None, + summary_scope: str = "pages", use_embedded_toc: bool = True, optimize: str = "full"): self._store = DocStore(storage_path) @@ -49,6 +51,8 @@ def __init__(self, storage_path: str, model: str, summary_model: str, self._index_backend = index_backend self._summary_max_words = summary_max_words self._summary_concurrency = summary_concurrency + self._summary_max_input_tokens = summary_max_input_tokens + self._summary_scope = summary_scope self._use_embedded_toc = use_embedded_toc self._optimize = optimize from .utils import ConfigLoader @@ -107,11 +111,12 @@ def submit_document( if mode is None: mode = "flash" if mode == "standard" and (self._summary_max_words is not None - or self._summary_concurrency is not None): + or self._summary_concurrency is not None + or self._summary_scope != "pages"): raise PageIndexAPIError( - "Failed to submit document: summary_max_words and " - "summary_concurrency are flash-only; mode='standard' does not " - "support them.") + "Failed to submit document: summary_max_words, " + "summary_concurrency and summary_scope are flash-only; " + "mode='standard' does not support them.") file_path = os.path.abspath(os.path.expanduser(str(file_path))) if not os.path.isfile(file_path): raise FileNotFoundError(f"No such file: {file_path}") @@ -245,6 +250,8 @@ def _index_flash(self, file_path: str) -> tuple[list, str | None]: optimize_model=self._summary_model, summary_concurrency=self._summary_concurrency, summary_max_words=self._summary_max_words, + summary_max_input_tokens=self._summary_max_input_tokens, + summary_scope=self._summary_scope, use_embedded_toc=self._use_embedded_toc) structure = result.get("structure", []) reason = flash_rejection_reason(result) @@ -254,6 +261,7 @@ def _index_flash(self, file_path: str) -> tuple[list, str | None]: description = generate_doc_description( create_clean_structure_for_description(structure), model=self._summary_model, + max_input_tokens=self._summary_max_input_tokens, ) return structure, description diff --git a/pageindex/tree_optimize.py b/pageindex/tree_optimize.py index e607d5f00..401e30d18 100644 --- a/pageindex/tree_optimize.py +++ b/pageindex/tree_optimize.py @@ -57,13 +57,14 @@ import asyncio import copy import json +import logging import os import re import sys from types import SimpleNamespace -from .utils import (ConfigLoader, _is_unrecoverable, intro_title, is_intro, - llm_acompletion, strip_internal_keys) +from .utils import (ConfigLoader, _is_unrecoverable, budget_tokens, input_budget, + intro_title, is_intro, llm_acompletion, strip_internal_keys) TRIGGER_PAGES = 5 # only look ahead on nodes larger than this ROUTING_COST = 1 # R(v), in pages @@ -223,8 +224,11 @@ def add_intro_nodes(structure, lines=None): opens = bool(lines) and first <= len(lines) and heading_at_page_start( lines, first, children[0]["title"]) end = max(node["start_index"], min(node["end_index"], first - 1 if opens else first)) - node["nodes"] = [{"title": intro_title(node.get("title")), - "start_index": node["start_index"], "end_index": end}] + children + intro = {"title": intro_title(node.get("title")), + "start_index": node["start_index"], "end_index": end} + if node.get("_pos") is not None: # the intro's text starts at the parent's heading + intro["_pos"] = node["_pos"] + node["nodes"] = [intro] + children return structure @@ -656,6 +660,11 @@ async def propose_children(node, pages, args): return [] # the whole span is beyond the loaded pages block = "\n".join( f"\n{pages[n - 1][:PAGE_CHARS]}\n" for n in range(start, end + 1)) + budget = getattr(args, "input_budget", None) + if budget is not None and budget_tokens(block, model=args.model) > budget: + logging.warning("expand skipped for %r: pages %d-%d exceed the %d token input budget", + node["title"], start, end, budget) + return [] answer = await ask_model(args.model, EXPAND_PROMPT.format( title=node["title"], start=start, end=end, pages=block)) @@ -816,7 +825,7 @@ async def optimize(structure, pages, lines, model=None, routing=ROUTING_COST, do_merge=True, do_expand=True, max_rounds=3, page_count=None, cache=None, kinds=("section", "table"), empty_retries=1, do_relabel=True, progress=False, on_final=None, - concurrency=None): + concurrency=None, max_input_tokens=None): """Run merge and expand over a tree until neither changes anything. Mutates `structure` in place and returns a summary. @@ -844,6 +853,7 @@ def settled(nodes): kinds=set(kinds) if kinds else None, empty_retries=empty_retries, progress=progress, settled=settled, do_merge=do_merge, + input_budget=input_budget(max_input_tokens), concurrency=min(EXPAND_CONCURRENCY, concurrency or EXPAND_CONCURRENCY)) baseline = set(validate(structure, page_count)) if page_count else set() diff --git a/pageindex/types.py b/pageindex/types.py index 68e9ea54f..e8f77e581 100644 --- a/pageindex/types.py +++ b/pageindex/types.py @@ -34,6 +34,8 @@ class LocalIndexConfig(TypedDict, total=False): summary_model: str summary_max_words: int summary_concurrency: int + summary_max_input_tokens: int + summary_scope: Literal["pages", "section"] use_embedded_toc: bool optimize: Literal["full", "merge", "off"] backend: dict diff --git a/pageindex/utils.py b/pageindex/utils.py index 4981a63c5..9e01768fa 100644 --- a/pageindex/utils.py +++ b/pageindex/utils.py @@ -20,6 +20,7 @@ import yaml from pathlib import Path from types import SimpleNamespace as config +import math import re # litellm is imported inside the functions that use it; eager import is slow @@ -820,6 +821,40 @@ async def generate_summaries_for_structure(structure, model=None): SUMMARY_RAW_TEXT_TOKENS = 200 # leaves under this reuse their raw text as the summary SUMMARY_INTRO_MAX_PAGES = 3 # cap on leading pages fed into a parent summary SUMMARY_MAX_WORDS = 150 # word cap the summary prompts ask for +INPUT_BUDGET_MARGIN = 0.85 # share of max_input_tokens a prompt's text may use +INPUT_BUDGET_OVERHEAD = 200 # tokens kept for the instructions and the reply + + +def input_budget(max_input_tokens): + """Tokens of document text one prompt may carry, or None when unbounded.""" + if max_input_tokens is None: + return None + return max(int(max_input_tokens * INPUT_BUDGET_MARGIN) - INPUT_BUDGET_OVERHEAD, 1) + + +def budget_tokens(text, model=None): + """count_tokens, plus the digits a per-digit tokenizer counts and tiktoken packs three to a token.""" + return count_tokens(text, model=model) + sum(len(d) - math.ceil(len(d) / 3) for d in re.findall(r"\d+", text)) + + +def _pack(units, budget, model=None): + """Whole units grouped to at most `budget` tokens; a unit over budget is cut at spaces.""" + flat = [] + for unit in units: + if budget_tokens(unit, model=model) > budget and " " in unit: + flat += [" ".join(group) for group in _pack(unit.split(" "), budget, model)] + else: + flat.append(unit) + groups, size = [], 0 + for unit in flat: + tokens = budget_tokens(unit, model=model) + 1 + if groups and size + tokens <= budget: + groups[-1].append(unit) + size += tokens + else: + groups.append([unit]) + size = tokens + return groups class _PriorityGate: @@ -899,6 +934,44 @@ def _reply_json(reply): return None +_JSON_ESCAPES = {'n': '\n', 't': '\t', 'r': '\r', 'b': '\b', 'f': '\f'} + + +def _decode_escapes(raw): + """Decode the escapes of a JSON string written by a model that does not escape + LaTeX: inside $...$ a backslash before a letter starts a command (\\nu, \\times), + and so does \\t, \\r, \\b or \\f before a lowercase letter outside it (\\text, \\ref).""" + def decode(text, math): + def one(m): + c = m.group(1) + if len(c) == 5: + return chr(int(c[1:], 16)) + if c in '"\\/': + return c + latex = math or (c in 'trbf' and text[m.end():m.end() + 1].islower()) + return _JSON_ESCAPES[c] if c in _JSON_ESCAPES and not latex else m.group(0) + return re.sub(r'\\(u[0-9a-fA-F]{4}|.)', one, text, flags=re.S) + return ''.join(decode(part, i % 2) for i, part in enumerate(re.split(r'(\$[^$]*\$)', raw))) + + +def _mangled(text): + """True when JSON decoding turned LaTeX commands into control characters.""" + return bool(re.search(r'[\b\t\f\r]', text)) + + +def _unparsed_field(reply, key): + """The value of `key` read from the raw reply, or None.""" + match = isinstance(reply, str) and re.search( + rf'"{key}"\s*:\s*"(.*?)"\s*(?:,\s*"[^"]+"\s*:|\}}\s*(?:```)?\s*$)', reply.strip(), re.S) + return _decode_escapes(match.group(1)) if match else None + + +def _field(reply, parsed, key): + """`key` of the parsed reply, read raw instead when decoding mangled its LaTeX.""" + value = parsed.get(key) + return (_unparsed_field(reply, key) or value) if isinstance(value, str) and _mangled(value) else value + + def parse_summary(reply): """The `summary` field of a model reply, or the reply itself when there is no such field.""" @@ -906,11 +979,11 @@ def parse_summary(reply): return "" parsed = _reply_json(reply) if isinstance(parsed, dict) and 'summary' in parsed: - summary = parsed['summary'] + summary = _field(reply, parsed, 'summary') if isinstance(summary, list): summary = ' '.join(str(item).strip() for item in summary if str(item).strip()) return str(summary).strip() if summary else "" - return reply.strip() + return (_unparsed_field(reply, 'summary') or reply).strip() def parse_title(reply): @@ -921,9 +994,7 @@ def parse_title(reply): deterministic one it already has. """ parsed = _reply_json(reply) - if not isinstance(parsed, dict): - return "" - title = parsed.get('title') + title = _field(reply, parsed, 'title') if isinstance(parsed, dict) else _unparsed_field(reply, 'title') if isinstance(title, list): title = ' '.join(str(item).strip() for item in title if str(item).strip()) return ' '.join(str(title).split()) if title else "" @@ -936,6 +1007,7 @@ def strip_internal_keys(structure): if not isinstance(node, dict): continue node.pop('_same_page', None) + node.pop('_pos', None) if node.get('nodes'): strip_internal_keys(node['nodes']) return structure @@ -960,13 +1032,15 @@ class SummaryScheduler: def __init__(self, structure, pdf_pages, model=None, small_node_tokens=SUMMARY_RAW_TEXT_TOKENS, max_intro_pages=SUMMARY_INTRO_MAX_PAGES, concurrency=None, - max_words=None): + max_words=None, max_input_tokens=None, blocks=None): self.structure = structure self._pdf_pages = pdf_pages self._model = model self._small_node_tokens = small_node_tokens self._max_intro_pages = max_intro_pages self._max_words = max_words or SUMMARY_MAX_WORDS + self._budget = input_budget(max_input_tokens) + self._blocks = blocks self._gate = _PriorityGate(concurrency or SUMMARY_CONCURRENCY) self._asked = self._answered = False self._marks = {} # id(node) -> future resolved once the node is final @@ -1015,9 +1089,26 @@ async def _ask(self, prompt, prio): self._answered = True return reply + def _section_text(self, node): + """The blocks from this node's heading to the next heading in the document, or None.""" + start = node.get('_pos') + if self._blocks is None or start is None: + return None + nodes = list(_subtree(self.structure)) + if any(n is not node and n.get('_pos') is None + and node['start_index'] <= n['start_index'] <= node['end_index'] + and not any(m is node for m in _subtree(n.get('nodes') or [])) + for n in nodes): + return None # a section without a located heading may start inside this one + later = [p for p in (n.get('_pos') for n in nodes) if p is not None and p > start] + return "\n".join(self._blocks[start:min(later, default=len(self._blocks))]) + async def _leaf_summary(self, node, prio): - text = get_text_of_pdf_pages(self._pdf_pages, node['start_index'], node['end_index']) - if count_tokens(text, model=self._model) < self._small_node_tokens: + text = self._section_text(node) + if text is None: + text = get_text_of_pdf_pages(self._pdf_pages, node['start_index'], node['end_index']) + tokens = count_tokens(text, model=self._model) + if tokens < self._small_node_tokens: return text.strip() # A node merged from same-page siblings carries a title joined from theirs. @@ -1032,7 +1123,8 @@ async def _leaf_summary(self, node, prio): title_field = ('\n "title": ,' if retitle else "") - prompt = f"""You are given a text chunk from a document. + def prompt_for(text): + return f"""You are given a text chunk from a document. Your task is to generate a concise description of everything that is covered in the text, summarizing all its points without omitting any type of content. Keep the description concise and to the point, avoiding unnecessary details, within {self._max_words} words.{ask_title} @@ -1045,13 +1137,43 @@ async def _leaf_summary(self, node, prio): Follow strictly the above JSON return format. Do not include any other text! """ - reply = await self._ask(prompt, prio) + if self._budget is not None and not retitle and budget_tokens(text, self._model) > self._budget: + chunks = ["\n".join(lines) for lines in _pack(text.split("\n"), self._budget, self._model)] + replies = await asyncio.gather(*(self._ask(prompt_for(chunk), prio) for chunk in chunks)) + return await self._reduce([parse_summary(reply) for reply in replies], prio) + reply = await self._ask(prompt_for(text), prio) if retitle: written = parse_title(reply) if written: node['title'] = written return parse_summary(reply) + async def _combine(self, summaries, prio): + parts = "\n".join(f"Part {i}: {summary}" for i, summary in enumerate(summaries, 1)) + prompt = f"""You are given summaries of consecutive parts of one section of a document. + Your task is to combine them into a single concise description of the whole section, within {self._max_words} words. + + Part Summaries: {parts} + + Reply strictly in the following JSON format: + {{ + "summary": + }} + + Follow strictly the above JSON return format. Do not include any other text! + """ + return parse_summary(await self._ask(prompt, prio)) + + async def _reduce(self, summaries, prio): + """Combine part summaries in budget-sized groups, level by level, into one.""" + while True: + groups = _pack(summaries, self._budget, self._model) + if len(groups) == len(summaries): # no two fit together: pair them anyway + groups = [summaries[i:i + 2] for i in range(0, len(summaries), 2)] + summaries = await asyncio.gather(*(self._combine(group, prio) for group in groups)) + if len(summaries) == 1: + return summaries[0] + async def _parent_summary(self, node, prio): children = node['nodes'] intro = get_intro_text(node, self._pdf_pages, max_pages=self._max_intro_pages) @@ -1140,8 +1262,8 @@ def _any_summary(nodes): async def summarize_tree(structure, pdf_pages, model=None, small_node_tokens=SUMMARY_RAW_TEXT_TOKENS, max_intro_pages=SUMMARY_INTRO_MAX_PAGES, concurrency=None, - max_words=None): - """Bottom-up summaries: leaves from their own pages, parents composed from + max_words=None, max_input_tokens=None, blocks=None): + """Bottom-up summaries: leaves from their own pages (or section blocks), parents composed from child summaries plus the pages no child covers. A parent's summary describes its whole subtree (end_index union semantics). Nodes that already carry a summary are left untouched; leaves under `small_node_tokens` use their raw @@ -1152,7 +1274,8 @@ async def summarize_tree(structure, pdf_pages, model=None, scheduler = SummaryScheduler(structure, pdf_pages, model=model, small_node_tokens=small_node_tokens, max_intro_pages=max_intro_pages, - concurrency=concurrency, max_words=max_words) + concurrency=concurrency, max_words=max_words, + max_input_tokens=max_input_tokens, blocks=blocks) scheduler.mark_final(list(_subtree(structure))) return await scheduler.finish() @@ -1180,7 +1303,33 @@ def create_clean_structure_for_description(structure): return structure -def generate_doc_description(structure, model=None): +def _tree_depth(nodes): + return 1 + max((_tree_depth(n["nodes"]) for n in nodes if n.get("nodes")), default=0) + + +def _prune_depth(nodes, depth): + return [{**{k: v for k, v in n.items() if k != "nodes"}, + **({"nodes": _prune_depth(n["nodes"], depth - 1)} if depth and n.get("nodes") else {})} + for n in nodes] + + +def fit_structure(structure, budget, model=None): + """The structure, cut from its deepest level up until it fits `budget` tokens.""" + if budget_tokens(str(structure), model=model) <= budget: + return structure + pruned = structure + for depth in range(_tree_depth(structure) - 2, -1, -1): + pruned = _prune_depth(structure, depth) + if budget_tokens(str(pruned), model=model) <= budget: + return pruned + logging.warning("document structure exceeds the %d token input budget", budget) + return pruned + + +def generate_doc_description(structure, model=None, max_input_tokens=None): + budget = input_budget(max_input_tokens) + if budget is not None: + structure = fit_structure(structure, budget, model) prompt = f"""Your are an expert in generating descriptions for a document. You are given a structure of a document. Your task is to generate a one-sentence description for the document, which makes it easy to distinguish the document from other documents. diff --git a/tests/test_client.py b/tests/test_client.py index 814415ca6..9fa3cb82d 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -3359,8 +3359,9 @@ def test_standard_mode_refuses_the_flash_summary_knobs(tmp_path, sample_pdf, mon monkeypatch.setattr(classic, "page_index_main", lambda *a, **kw: { "structure": [{"title": "T", "start_index": 1, "end_index": 1, "nodes": []}]}) monkeypatch.chdir(tmp_path) - for name, value in (("summary_concurrency", 2), ("summary_max_words", 7)): - with pytest.raises(PageIndexAPIError, match="summary_concurrency are flash-only"): + for name, value in (("summary_concurrency", 2), ("summary_max_words", 7), + ("summary_scope", "section")): + with pytest.raises(PageIndexAPIError, match="summary_scope are flash-only"): PageIndexLocalClient(**{name: value}).submit_document(sample_pdf, mode="standard") PageIndexLocalClient().submit_document(sample_pdf, mode="standard") diff --git a/tests/test_summary_budget.py b/tests/test_summary_budget.py new file mode 100644 index 000000000..1901d23ba --- /dev/null +++ b/tests/test_summary_budget.py @@ -0,0 +1,208 @@ +"""summary_max_input_tokens: prompts that grow with the document stay within the model's context.""" +import logging + +import pytest + +import pageindex.utils as utils +from pageindex import PageIndexLocalClient +from pageindex.client import PageIndexAPIError + + +def _tree(): + long = "word " * 400 + return [{"title": "A", "node_id": "0000", "summary": long, "nodes": [ + {"title": "A1", "node_id": "0001", "summary": long, "nodes": [ + {"title": "A1a", "node_id": "0002", "summary": long}]}]}, + {"title": "B", "node_id": "0003", "summary": long}] + + +def _titles(nodes): + return [n["title"] for n in nodes] + [t for n in nodes for t in _titles(n.get("nodes", []))] + + +def test_input_budget(): + assert utils.input_budget(None) is None + assert utils.input_budget(8192) == 6763 + assert utils.input_budget(10) == 1 + + +def test_fit_structure_returns_a_structure_that_fits_unchanged(): + tree = _tree() + assert utils.fit_structure(tree, 10 ** 6) is tree + + +def test_fit_structure_cuts_from_the_deepest_level(): + tree = _tree() + full = utils.count_tokens(str(tree)) + two_levels = utils.count_tokens(str(utils._prune_depth(tree, 1))) + one_level = utils.count_tokens(str(utils._prune_depth(tree, 0))) + assert one_level < two_levels < full + assert _titles(utils.fit_structure(tree, full - 1)) == ["A", "B", "A1"] + assert _titles(utils.fit_structure(tree, two_levels - 1)) == ["A", "B"] + assert tree == _tree() + + +def test_fit_structure_warns_when_the_top_level_alone_is_too_big(caplog): + with caplog.at_level(logging.WARNING): + out = utils.fit_structure(_tree(), 10) + assert _titles(out) == ["A", "B"] + assert "exceeds the 10 token input budget" in caplog.text + + +def _description_prompt(monkeypatch, **kwargs): + prompts = [] + monkeypatch.setattr(utils, "llm_completion", lambda model, prompt: prompts.append(prompt) or "d") + utils.generate_doc_description(_tree(), **kwargs) + return prompts[0] + + +def test_description_prompt_is_unchanged_without_a_budget(monkeypatch): + assert f"Document Structure: {_tree()}" in _description_prompt(monkeypatch) + + +def test_description_prompt_is_cut_to_the_budget(monkeypatch): + prompt = _description_prompt(monkeypatch, max_input_tokens=1000) + assert "'title': 'A'" in prompt and "'title': 'A1a'" not in prompt + + +def test_client_takes_the_budget_flat_or_in_the_index_slot(): + assert PageIndexLocalClient(summary_max_input_tokens=8192)._api._summary_max_input_tokens == 8192 + assert PageIndexLocalClient(index={"summary_max_input_tokens": 8192})._api._summary_max_input_tokens == 8192 + assert PageIndexLocalClient()._api._summary_max_input_tokens is None + with pytest.raises(PageIndexAPIError, match="summary_max_input_tokens must be a positive int"): + PageIndexLocalClient(summary_max_input_tokens=-1) + with pytest.raises(PageIndexAPIError, match="summary_max_input_tokens must be a"): + PageIndexLocalClient(summary_max_input_tokens="8192") + + +def test_the_budget_reaches_the_description(tmp_path, sample_pdf, monkeypatch): + import pageindex.flash + seen = {} + monkeypatch.setattr(pageindex.flash, "page_index_flash", lambda p, **kw: { + "structure": [{"title": "T", "start_index": 1, "end_index": 1, "summary": "s"}]}) + monkeypatch.setattr(utils, "generate_doc_description", + lambda structure, model=None, max_input_tokens=None: + seen.update(max_input_tokens=max_input_tokens) or "d") + monkeypatch.chdir(tmp_path) + PageIndexLocalClient(summary_max_input_tokens=8192).submit_document(sample_pdf) + assert seen == {"max_input_tokens": 8192} + + +def _leaf(monkeypatch, lines, max_input_tokens, reply=lambda prompt: '{"summary": "ok"}', **node): + import asyncio + prompts = [] + + async def fake(model, prompt): + prompts.append(prompt) + return reply(prompt) + monkeypatch.setattr(utils, "llm_acompletion", fake) + structure = [{"title": "T", "start_index": 1, "end_index": 1, **node}] + asyncio.run(utils.summarize_tree(structure, [("\n".join(lines), 0)], small_node_tokens=0, + max_input_tokens=max_input_tokens)) + return prompts, structure[0]["summary"] + + +def _chunk(prompt): + return prompt.split("Given Text: ")[1].split("\n\n Reply strictly")[0].split("\n") + + +def test_leaf_within_the_budget_is_one_call(monkeypatch): + lines = [f"line {i} about apples" for i in range(20)] + assert len(_leaf(monkeypatch, lines, None)[0]) == 1 + assert len(_leaf(monkeypatch, lines, 8192)[0]) == 1 + + +def test_long_leaf_is_split_by_lines_and_combined(monkeypatch): + lines = [f"line {i} " + "word " * 10 for i in range(400)] + budget = utils.input_budget(1200) + prompts, summary = _leaf(monkeypatch, lines, 1200, + reply=lambda p: '{"summary": "combined"}' if "Part Summaries" in p else '{"summary": "part"}') + parts, combine = prompts[:-1], prompts[-1] + assert len(parts) > 2 and "Part Summaries" in combine and summary == "combined" + assert [line for p in parts for line in _chunk(p)] == lines + assert all(utils.count_tokens("\n".join(_chunk(p))) <= budget for p in parts) + assert combine.count("Part ") == len(parts) + 1 # one per part, plus the instruction + + +def test_many_parts_are_combined_in_levels(monkeypatch): + lines = ["word " * 40 for _ in range(300)] + prompts, summary = _leaf(monkeypatch, lines, 1000, + reply=lambda p: '{"summary": "' + "long " * 120 + '"}') + assert sum("Part Summaries" in p for p in prompts) > 1 + assert summary.startswith("long") + + +def test_a_line_over_the_budget_is_cut_at_spaces(monkeypatch): + prompts, _ = _leaf(monkeypatch, ["start", "word " * 3000, "end"], 1000) + budget = utils.input_budget(1000) + parts = [p for p in prompts if "Part Summaries" not in p] + assert len(parts) > 3 + assert all(utils.count_tokens("\n".join(_chunk(p))) <= budget for p in parts) + assert _chunk(parts[0])[0] == "start" and _chunk(parts[-1])[-1].endswith("end") + + +def test_a_node_merged_from_same_page_siblings_is_not_split(monkeypatch): + lines = [f"line {i} " + "word " * 10 for i in range(400)] + prompts, _ = _leaf(monkeypatch, lines, 1200, _same_page=True, key_items=["A", "B"]) + assert len(prompts) == 1 + + +def _expand(monkeypatch, args, pages): + import asyncio + from types import SimpleNamespace + import pageindex.tree_optimize as tree_optimize + asked = [] + + async def fake_ask(model, prompt): + asked.append(prompt) + return {"subsections": []} + monkeypatch.setattr(tree_optimize, "ask_model", fake_ask) + node = {"title": "References", "start_index": 1, "end_index": 3, "node_id": "n1"} + out = asyncio.run(tree_optimize.propose_children(node, pages, SimpleNamespace(model="m", **args))) + return out, asked + + +def test_expand_is_skipped_when_the_pages_overrun_the_budget(monkeypatch, caplog): + pages = ["word " * 400] * 3 + with caplog.at_level(logging.WARNING): + out, asked = _expand(monkeypatch, {"input_budget": 500}, pages) + assert out == [] and asked == [] + assert "expand skipped for 'References': pages 1-3 exceed the 500 token input budget" in caplog.text + + +def test_expand_asks_the_model_within_the_budget_or_without_one(monkeypatch): + pages = ["word " * 400] * 3 + assert len(_expand(monkeypatch, {"input_budget": 10 ** 6}, pages)[1]) == 1 + assert len(_expand(monkeypatch, {"input_budget": None}, pages)[1]) == 1 + assert len(_expand(monkeypatch, {}, pages)[1]) == 1 + + +def test_optimize_leaves_the_node_collapsed_when_the_budget_is_too_small(monkeypatch): + import asyncio + import pageindex.tree_optimize as tree_optimize + from test_client import _expand_fixture + tree, pages, lines = _expand_fixture() + asked = [] + + async def fake_ask(model, prompt): + asked.append(prompt) + return {"subsections": [{"title": "Sub One", "page": 4}]} + monkeypatch.setattr(tree_optimize, "ask_model", fake_ask) + asyncio.run(tree_optimize.optimize(tree, pages, lines, model="m", do_expand=True, + max_input_tokens=300)) + assert asked == [] + + +def test_budget_tokens_count_the_digits_tiktoken_packs_together(): + text = "kernel, 8, 10, 317, 3482, 2024" + assert utils.budget_tokens(text) == utils.count_tokens(text) + 0 + 1 + 2 + 2 + 2 # 8, 10, 317, 3482, 2024 + assert utils.budget_tokens("no digits here") == utils.count_tokens("no digits here") + + +def test_an_index_full_of_page_numbers_is_split_by_its_digit_aware_size(monkeypatch): + lines = [f"term {i}, " + ", ".join(str(100 + 7 * k) for k in range(30)) for i in range(120)] + budget = utils.input_budget(1200) + prompts, _ = _leaf(monkeypatch, lines, 1200) + parts = [p for p in prompts if "Part Summaries" not in p] + assert all(utils.budget_tokens("\n".join(_chunk(p))) <= budget for p in parts) + assert any(utils.count_tokens("\n".join(_chunk(p))) < budget * 0.8 for p in parts) diff --git a/tests/test_summary_replies.py b/tests/test_summary_replies.py new file mode 100644 index 000000000..362a2d185 --- /dev/null +++ b/tests/test_summary_replies.py @@ -0,0 +1,53 @@ +"""Summary replies whose JSON does not parse.""" +import pytest + +import pageindex.utils as utils + + +@pytest.mark.parametrize("reply, expected", [ + # unescaped quotes inside the text + ('{\n "summary": "A new "herd" of models."\n}', 'A new "herd" of models.'), + # stray backslash from LaTeX, inside a fence + ('```json\n{\n "summary": "Scaling $N^*(C)=AC^\\alpha$ holds."\n}\n```', + 'Scaling $N^*(C)=AC^\\alpha$ holds.'), + # escaped quote next to an unescaped one + ('{"summary": "He said \\"yes\\" and "no"."}', 'He said "yes" and "no".'), + # valid JSON, prose and a missing field behave as before + ('{"summary": "Plain."}', "Plain."), + ("Just a sentence.", "Just a sentence."), + ('{"points": []}', '{"points": []}'), +]) +def test_parse_summary(reply, expected): + assert utils.parse_summary(reply) == expected + + +def test_parse_title_reads_the_field_of_unparseable_json(): + reply = '{"title": "A "big" page", "summary": "Text."}' + assert utils.parse_title(reply) == 'A "big" page' + assert utils.parse_summary(reply) == "Text." + + +def test_parse_title_still_refuses_a_reply_without_the_field(): + assert utils.parse_title("Just a sentence.") == "" + assert utils.parse_title(None) == "" + + +def test_latex_in_invalid_json_keeps_its_backslashes(): + reply = r'{"summary": "Shows $\alpha \to \nu$ as $L\times W$ \u2190 width.\n\nThen \"more\"."}' + assert utils.parse_summary(reply) == 'Shows $\\alpha \\to \\nu$ as $L\\times W$ ← width.\n\nThen "more".' + + +def test_latex_that_json_decodes_to_control_characters_is_read_raw(): + reply = r'{"summary": "A $W \times L$ grid of $\boldsymbol{x}$ and $\nu$.", "title": "The $\times$ map"}' + assert utils.parse_summary(reply) == r"A $W \times L$ grid of $\boldsymbol{x}$ and $\nu$." + assert utils.parse_title(reply) == r"The $\times$ map" + + +def test_valid_json_with_escaped_latex_and_paragraphs_is_unchanged(): + assert utils.parse_summary(r'{"summary": "First $\\alpha$.\n\nSecond."}') == "First $\\alpha$.\n\nSecond." + assert utils.parse_summary(r'{"summary": "Costs $5.\n\nLater $10."}') == "Costs $5.\n\nLater $10." + + +def test_latex_commands_outside_math_are_kept(): + reply = r'{"summary": "A leap (\text{Eq. } \ref{eq:71}).\n\nThen\tDone."}' + assert utils.parse_summary(reply) == 'A leap (\\text{Eq. } \\ref{eq:71}).\n\nThen\tDone.' diff --git a/tests/test_summary_scope.py b/tests/test_summary_scope.py new file mode 100644 index 000000000..4eb2f225d --- /dev/null +++ b/tests/test_summary_scope.py @@ -0,0 +1,171 @@ +"""summary_scope="section": a leaf is summarized from the blocks between its heading and the next one.""" +import asyncio +from pathlib import Path + +import pytest + +import pageindex.utils as utils +from pageindex import PageIndexLocalClient +from pageindex.client import PageIndexAPIError +from pageindex.flash import page_index_flash +from pageindex.flash.main import extract_toc + +PDF = str(Path(__file__).parent / "data" / "flash" / "ja_report.pdf") + + +def _flat(nodes): + for n in nodes: + yield n + yield from _flat(n.get("nodes", [])) + + +def test_extract_toc_reports_the_heading_blocks_on_request(): + result = extract_toc(PDF, use_embedded_toc=False, with_blocks=True) + nodes = list(_flat(result["structure"])) + positions = [n["_pos"] for n in nodes] + assert nodes and all(isinstance(p, int) for p in positions) and positions == sorted(positions) + assert all(n["title"] in result["block_texts"][n["_pos"]] for n in nodes) + + +def test_extract_toc_is_unchanged_without_blocks(): + result = extract_toc(PDF, use_embedded_toc=False) + assert "block_texts" not in result + assert not any("_pos" in n for n in _flat(result["structure"])) + + +def test_the_internal_positions_do_not_reach_the_result(): + result = page_index_flash(PDF, summary=False, optimize=False, use_embedded_toc=False, + summary_scope="section") + assert "block_texts" not in result + assert not any("_pos" in n for n in _flat(result["structure"])) + + +def _prompts(structure, blocks, pages=(("page one\npage two", 0),)): + prompts = [] + + async def fake(model, prompt): + prompts.append(prompt.split("Given Text: ")[1].split("\n\n Reply strictly")[0]) + return '{"summary": "ok"}' + utils.llm_acompletion, saved = fake, utils.llm_acompletion + try: + asyncio.run(utils.summarize_tree(structure, list(pages), small_node_tokens=0, blocks=blocks)) + finally: + utils.llm_acompletion = saved + return prompts + + +def _tree(*positions): + return [{"title": f"S{i}", "start_index": 1, "end_index": 1, **({} if p is None else {"_pos": p})} + for i, p in enumerate(positions)] + + +def test_a_leaf_is_summarized_from_its_own_blocks(): + prompts = _prompts(_tree(0, 3), ["a0", "a1", "a2", "b0", "b1"]) + assert sorted(prompts) == ["a0\na1\na2", "b0\nb1"] + + +def test_a_section_ends_at_the_next_heading_in_the_document_not_in_the_tree(): + prompts = _prompts(_tree(3, 0), ["a0", "a1", "a2", "b0", "b1"]) + assert sorted(prompts) == ["a0\na1\na2", "b0\nb1"] + assert _prompts(_tree(0), ["a0", "a1"]) == ["a0\na1"] + + +def test_a_leaf_falls_back_to_its_pages_without_its_own_position_or_blocks(): + blocks = ["a0", "a1", "b0"] + assert _prompts(_tree(None, 2), blocks)[0] == "page one\npage two" + assert _prompts(_tree(0, 2), None) == ["page one\npage two"] * 2 + + +def test_a_leaf_falls_back_to_its_pages_when_a_section_without_a_position_starts_inside_it(): + pages = (("page one", 0), ("page two", 0), ("page three", 0)) + blocks = ["a0", "a1", "b0"] + + def tree(unlocated_page): + return [{"title": "S0", "start_index": 1, "end_index": 2, "_pos": 0}, + {"title": "S1", "start_index": unlocated_page, "end_index": unlocated_page}, + {"title": "S2", "start_index": 3, "end_index": 3, "_pos": 2}] + assert set(_prompts(tree(2), blocks, pages)) == {"page onepage two", "page two", "b0"} + assert "a0\na1" in set(_prompts(tree(3), blocks, pages)) + + +def test_the_internal_position_is_dropped_from_the_summarized_tree(): + structure = _tree(0, 2) + _prompts(structure, ["a0", "a1", "b0"]) + assert not any("_pos" in n for n in structure) + + +def test_client_takes_the_scope_flat_or_in_the_index_slot(): + assert PageIndexLocalClient(summary_scope="section")._api._summary_scope == "section" + assert PageIndexLocalClient(index={"summary_scope": "section"})._api._summary_scope == "section" + assert PageIndexLocalClient()._api._summary_scope == "pages" + with pytest.raises(PageIndexAPIError, match='summary_scope must be "pages" or "section"'): + PageIndexLocalClient(summary_scope="paragraph") + + +def test_flash_refuses_an_unknown_scope(): + with pytest.raises(ValueError, match='summary_scope must be "pages" or "section"'): + page_index_flash(PDF, summary_scope="paragraph") + + +def test_the_scope_reaches_the_flash_indexer(tmp_path, sample_pdf, monkeypatch): + import pageindex.flash + seen = {} + monkeypatch.setattr(pageindex.flash, "page_index_flash", lambda p, **kw: seen.update(kw) or { + "structure": [{"title": "T", "start_index": 1, "end_index": 1, "summary": "s"}]}) + monkeypatch.setattr(utils, "llm_completion", lambda model, prompt, **kw: "d") + monkeypatch.chdir(tmp_path) + PageIndexLocalClient(summary_scope="section").submit_document(sample_pdf) + assert seen["summary_scope"] == "section" + + +def test_a_title_is_located_by_its_heading_block_then_by_a_block_that_is_mostly_the_title(): + from pageindex.flash.main import _find_heading + blocks = [(0, "mgsm is evaluated across seven languages and the results are listed below", 0), + (1, "5 2 4 multilingual benchmarks", 7), + (2, "mgsm", 0), + (3, "safety pretraining", 7)] + assert _find_heading("Multilingual Benchmarks", blocks) == 1 # numbering prefix + assert _find_heading("MGSM", blocks) == 2 # not the paragraph that starts with it + assert _find_heading("Safety Pre-training", blocks) == 3 # close enough + assert _find_heading("Something else", blocks) is None + + +def test_the_closest_match_on_a_page_wins_over_the_first_one(): + from pageindex.flash.main import _find_heading + blocks = [(0, "contributors and acknowledgements", 7), (1, "6 1 pre training evaluations", 7), + (2, "contributors", 7), (3, "6 2 post training evaluations", 7)] + assert _find_heading("Contributors", blocks) == 2 + assert _find_heading("6.2. Post-training Evaluations", blocks) == 3 + + +def test_a_bookmark_one_page_off_finds_the_heading_block_on_the_next_or_previous_page(): + from pageindex.flash.main import locate_headings + + class B: + def __init__(self, text, kind): + self.text, self.type = text, kind + import pageindex.flash.main as fm + body = [B("previous paragraph mentions reliability and operational challenges", 0), + B("3.3.4 Reliability and Operational Challenges", 8), B("body", 0), B("4 Annex", 7)] + saved, fm.block_text = fm.block_text, lambda b: b.text + try: + tree = [{"title": "Reliability and Operational Challenges", "start_index": 12}, + {"title": "Annex", "start_index": 14}, {"title": "Missing", "start_index": 12}] + locate_headings(tree, body, [12, 13, 13, 13]) + finally: + fm.block_text = saved + assert tree[0]["_pos"] == 1 and tree[1]["_pos"] == 3 and "_pos" not in tree[2] + + +def test_an_intro_node_starts_at_its_parent_heading(): + from pageindex.tree_optimize import add_intro_nodes + tree = [{"title": "Ch", "start_index": 1, "end_index": 3, "_pos": 0, + "nodes": [{"title": "Sec", "start_index": 2, "end_index": 3, "_pos": 2}]}, + {"title": "Next", "start_index": 3, "end_index": 3, "_pos": 4}] + add_intro_nodes(tree) + intro, sec = tree[0]["nodes"] + assert intro["title"] == "Ch (intro)" and intro["_pos"] == 0 + blocks = ["Ch", "chapter text", "Sec", "section text", "Next"] + prompts = _prompts(tree, blocks) + assert any("Ch\nchapter text" in p and "section text" not in p for p in prompts) + assert any("Sec\nsection text" in p and "Next" not in p for p in prompts)