Coverage for src / quber / core / consolidation.py: 65%
118 statements
« prev ^ index » next coverage.py v7.14.0, created at 2026-09-23 22:14 -0400
« prev ^ index » next coverage.py v7.14.0, created at 2026-09-23 22:14 -0400
1"""
2Header consolidation post-processor.
4Operates on a DoclingDocument (the canonical Parser output) and returns
5TableMetadata for each table. Pulled out of the parser to keep that layer
6engine-agnostic — the consolidation logic is independent of how the
7DoclingDocument was produced.
8"""
10import json
11from datetime import datetime
12from pathlib import Path
13from typing import Any, List, Optional, Tuple
15from docling_core.types.doc.document import DoclingDocument
17from quber.core.models import TableMetadata
20class HeaderConsolidator:
21 def __init__(
22 self,
23 separator: str = " ^ ",
24 max_lookback: int = 10,
25 enable_consolidation: bool = True,
26 debug: bool = False,
27 ) -> None:
28 self.separator = separator
29 self.max_lookback = max_lookback
30 self.enable_consolidation = enable_consolidation
31 self.debug = debug
33 def consolidate(self, document: DoclingDocument) -> List[TableMetadata]:
34 tables_metadata: List[TableMetadata] = []
35 table_index = 0
37 elements_with_depth = list(document.iterate_items())
38 elements_list = [item for item, _ in elements_with_depth]
39 element_to_index = {id(elem): i for i, elem in enumerate(elements_list)}
41 for table in document.tables:
42 page_num = (
43 table.prov[0].page_no
44 if (hasattr(table, "prov") and table.prov and hasattr(table.prov[0], "page_no"))
45 else 1
46 )
48 table_idx = element_to_index.get(id(table), -1)
50 headers: List[str] = []
51 descriptive_texts: List[str] = []
52 if table_idx > 0:
53 headers, descriptive_texts = self.find_preceding_headers(elements_list, table_idx)
55 consolidated_header: Optional[str] = None
56 if self.enable_consolidation and headers:
57 consolidated_header = self.consolidate_headers(headers)
59 rows = 0
60 cols = 0
61 first_cell: Optional[str] = None
62 if hasattr(table, "data") and table.data:
63 rows = table.data.num_rows
64 cols = table.data.num_cols
65 if rows > 0 and cols > 0:
66 if hasattr(table.data, "grid") and table.data.grid:
67 first_row = table.data.grid[0] if len(table.data.grid) > 0 else []
68 if first_row and len(first_row) > 0:
69 first_cell = str(first_row[0]) if first_row[0] else None
71 preceding_text = self.get_preceding_text(elements_list, table_idx, limit=3)
73 table_md = table.export_to_markdown(document) if hasattr(table, "export_to_markdown") else None
75 metadata = TableMetadata(
76 table_index=table_index,
77 page_number=page_num,
78 headers=headers,
79 consolidated_header=consolidated_header,
80 descriptive_text=descriptive_texts[0] if descriptive_texts else None,
81 table_markdown=table_md,
82 preceding_text=preceding_text,
83 rows=rows,
84 cols=cols,
85 first_cell_content=first_cell,
86 )
87 tables_metadata.append(metadata)
88 table_index += 1
90 if self.debug:
91 self.print_debug_info(metadata)
93 return tables_metadata
95 def find_preceding_headers(self, elements: List[Any], table_idx: int) -> Tuple[List[str], List[str]]:
96 headers: List[str] = []
97 descriptive_texts: List[str] = []
98 current_level: Optional[int] = None
99 consecutive = True
101 for i in range(min(self.max_lookback, table_idx)):
102 idx = table_idx - i - 1
103 elem = elements[idx]
104 elem_type = type(elem).__name__
106 if elem_type in ["TableItem", "PictureItem"]:
107 break
109 if elem_type == "ListItem" and headers:
110 break
112 if not hasattr(elem, "text") or not elem.text:
113 continue
115 text = elem.text.strip()
116 if not text:
117 continue
119 is_header = hasattr(elem, "level") and elem.level is not None
120 if is_header:
121 level = elem.level
122 if current_level is None:
123 current_level = level
124 headers.insert(0, text)
125 elif level == current_level and consecutive:
126 headers.insert(0, text)
127 else:
128 break
129 else:
130 if not headers and len(descriptive_texts) == 0:
131 descriptive_texts.insert(0, text)
132 elif headers:
133 consecutive = False
135 return headers, descriptive_texts
137 def consolidate_headers(self, headers: List[str]) -> Optional[str]:
138 if not headers:
139 return None
140 if len(headers) == 1:
141 return headers[0]
142 return self.separator.join(headers)
144 def get_preceding_text(self, elements: List[Any], table_idx: int, limit: int = 3) -> Optional[str]:
145 texts: List[str] = []
146 count = 0
148 for i in range(min(self.max_lookback, table_idx)):
149 idx = table_idx - i - 1
150 elem = elements[idx]
151 elem_type = type(elem).__name__
153 if elem_type == "TableItem":
154 texts.insert(0, "[Table]")
155 continue
157 if hasattr(elem, "text") and elem.text:
158 text = elem.text.strip()
159 if text:
160 texts.insert(0, text)
161 count += 1
162 if count >= limit:
163 break
165 return "\n".join(texts) if texts else None
167 def print_debug_info(self, metadata: TableMetadata) -> None:
168 print(f"\n=== Table {metadata.table_index + 1} (Page {metadata.page_number}) ===")
169 if metadata.consolidated_header:
170 print(f"Consolidated Header: {metadata.consolidated_header}")
171 if metadata.headers:
172 print(f"Individual Headers: {metadata.headers}")
173 if metadata.descriptive_text:
174 print(f"Descriptive Text: {metadata.descriptive_text[:100]}...")
175 print(f"Size: {metadata.rows} rows x {metadata.cols} cols")
176 if metadata.first_cell_content:
177 print(f"First Cell: {metadata.first_cell_content[:50]}...")
180def save_metadata(metadata_list: List[TableMetadata], output_path: str | Path) -> None:
181 data = {
182 "extraction_date": datetime.now().isoformat(),
183 "total_tables": len(metadata_list),
184 "tables": [
185 {
186 "table_index": meta.table_index,
187 "page_number": meta.page_number,
188 "consolidated_header": meta.consolidated_header,
189 "headers": meta.headers,
190 "descriptive_text": meta.descriptive_text,
191 "preceding_text": meta.preceding_text,
192 "dimensions": {"rows": meta.rows, "cols": meta.cols},
193 "first_cell": meta.first_cell_content,
194 }
195 for meta in metadata_list
196 ],
197 }
198 with open(output_path, "w", encoding="utf-8") as f:
199 json.dump(data, f, indent=2, ensure_ascii=False)