from io import BytesIO from copy import deepcopy from docx import Document from docx.document import Document as DocumentObject from docx.oxml import OxmlElement from docx.oxml.table import CT_Tbl from docx.oxml.text.paragraph import CT_P from docx.table import Table, _Cell from docx.text.paragraph import Paragraph def _iter_block_items(parent: DocumentObject | _Cell): parent_elm = parent.element.body if isinstance(parent, DocumentObject) else parent._tc for child in parent_elm.iterchildren(): if isinstance(child, CT_P): yield Paragraph(child, parent) elif isinstance(child, CT_Tbl): yield Table(child, parent) def _is_heading(paragraph: Paragraph) -> bool: style_name = paragraph.style.name if paragraph.style is not None else "" normalized = style_name.lower().replace(" ", "") return normalized.startswith("heading") def _delete_block(block): element = block._element parent = element.getparent() if parent is not None: parent.remove(element) def _delete_heading_section(heading: Paragraph): blocks = [heading] current = heading._element.getnext() while current is not None: if isinstance(current, CT_P): para = Paragraph(current, heading._parent) if _is_heading(para): break blocks.append(para) elif isinstance(current, CT_Tbl): blocks.append(Table(current, heading._parent)) current = current.getnext() for block in blocks: _delete_block(block) def _remove_unreferenced_headings(document: DocumentObject, referenced_anchors: set[str]): headings_to_remove: list[Paragraph] = [] found_first_heading = False pre_heading_blocks: list = [] print(f"[EXPORT] referenced_anchors: {referenced_anchors}") for block in _iter_block_items(document): if isinstance(block, Paragraph) and _is_heading(block): found_first_heading = True text = block.text.strip() if text not in referenced_anchors: print(f"[EXPORT] Unreferenced heading found, will remove: '{text}'") headings_to_remove.append(block) elif not found_first_heading: pre_heading_blocks.append(block) for heading in headings_to_remove: _delete_heading_section(heading) if not referenced_anchors: for block in pre_heading_blocks: _delete_block(block) def _clear_paragraph(paragraph: Paragraph): element = paragraph._element for child in list(element): if child.tag.endswith("}r"): element.remove(child) def _copy_paragraph_format(target: Paragraph, source: Paragraph | None): if source is None: return source_ppr = source._element.pPr if source_ppr is not None: target._element.insert(0, deepcopy(source_ppr)) def _copy_run_format(target_run, source_paragraph: Paragraph | None): if source_paragraph is None: return for source_run in source_paragraph.runs: if source_run._element.rPr is not None: target_run._element.insert(0, deepcopy(source_run._element.rPr)) break def _extract_first_run_format(source_paragraph: Paragraph | None): if source_paragraph is None: return None for source_run in source_paragraph.runs: if source_run._element.rPr is not None: return deepcopy(source_run._element.rPr) return None def _set_paragraph_text( paragraph: Paragraph, text: str, style_name: str | None = None, template_paragraph: Paragraph | None = None, ): run_format = _extract_first_run_format(template_paragraph) _clear_paragraph(paragraph) if style_name: try: paragraph.style = style_name except Exception: pass if text: run = paragraph.add_run(text) if run_format is not None: run._element.insert(0, run_format) def _append_paragraph_after( paragraph: Paragraph, text: str, style_name: str | None = None, template_paragraph: Paragraph | None = None, ) -> Paragraph: new_p = OxmlElement("w:p") paragraph._element.addnext(new_p) new_para = Paragraph(new_p, paragraph._parent) _copy_paragraph_format(new_para, template_paragraph) if style_name: try: new_para.style = style_name except Exception: pass if text: run = new_para.add_run(text) _copy_run_format(run, template_paragraph) return new_para def _set_cell_text_with_template(cell, value: str, template_paragraph: Paragraph | None = None): if not cell.paragraphs: cell.text = value return paragraph = cell.paragraphs[0] _clear_paragraph(paragraph) run = paragraph.add_run(value) _copy_run_format(run, template_paragraph) def _resize_table_rows(table: Table, row_count: int): current_rows = len(table.rows) if current_rows == 0: return if current_rows < row_count: template_row = table.rows[-1]._tr for _ in range(row_count - current_rows): table._tbl.append(deepcopy(template_row)) elif current_rows > row_count: for _ in range(current_rows - row_count): table._tbl.remove(table.rows[-1]._tr) def _fill_table(table: Table, matrix: list[list[str]]): if not matrix: return _resize_table_rows(table, len(matrix)) template_cell_paragraph = table.rows[0].cells[0].paragraphs[0] if table.rows and table.rows[0].cells else None for row_index, row_values in enumerate(matrix): row = table.rows[row_index] for col_index, cell in enumerate(row.cells): value = row_values[col_index] if col_index < len(row_values) else "" _set_cell_text_with_template(cell, value, template_cell_paragraph) def _append_table_after( paragraph: Paragraph, rows: list[list[str]], headers: list[str] | None = None, template_table: Table | None = None, ): matrix = [headers, *rows] if headers else rows if template_table is not None: cloned_tbl = deepcopy(template_table._tbl) paragraph._element.addnext(cloned_tbl) cloned_table = Table(cloned_tbl, paragraph._parent) _fill_table(cloned_table, matrix) return cloned_table container = paragraph._parent table = container.add_table(rows=max(len(matrix), 1), cols=max(len(headers or []), len(rows[0]) if rows else 1)) if headers: for row_index, row_values in enumerate(matrix): for index, value in enumerate(row_values): table.rows[row_index].cells[index].text = value elif rows: for row_index, row_values in enumerate(matrix): for index, value in enumerate(row_values): table.rows[row_index].cells[index].text = value tbl = table._tbl tbl.getparent().remove(tbl) paragraph._element.addnext(tbl) return Table(tbl, container) def _append_empty_paragraph_after_table(table: Table, style_name: str | None = None) -> Paragraph: new_p = OxmlElement("w:p") table._tbl.addnext(new_p) new_para = Paragraph(new_p, table._parent) if style_name: try: new_para.style = style_name except Exception: pass return new_para def _find_heading_paragraph(document: DocumentObject, heading_text: str, after_element=None) -> Paragraph | None: started = after_element is None for block in _iter_block_items(document): if isinstance(block, Paragraph) and _is_heading(block) and block.text.strip() == heading_text.strip(): if started: return block if after_element is not None and block._element == after_element: started = True return None def _collect_section_templates(heading: Paragraph): first_body_style = None paragraph_template = None table_template = None blocks = [] current = heading._element.getnext() while current is not None: if isinstance(current, CT_P): current_paragraph = Paragraph(current, heading._parent) if _is_heading(current_paragraph): break if first_body_style is None and current_paragraph.style is not None: first_body_style = current_paragraph.style.name if paragraph_template is None: paragraph_template = current_paragraph blocks.append(current_paragraph) elif isinstance(current, CT_Tbl): current_table = Table(current, heading._parent) if table_template is None: table_template = current_table blocks.append(current_table) current = current.getnext() return first_body_style, paragraph_template, table_template, blocks def _insert_content_after( insert_after: Paragraph, content: dict, first_body_style: str | None, paragraph_template: Paragraph | None, table_template: Table | None, ): current_anchor: Paragraph = insert_after content_blocks = content.get("content", []) for block in content_blocks: block_type = block.get("type") if block_type == "table": rows = [list(row) for row in block.get("rows", [])] headers = block.get("headers") or [] table = _append_table_after(current_anchor, rows, headers, table_template) current_anchor = _append_empty_paragraph_after_table(table, first_body_style) else: text = block.get("text", "") text_parts = [item for item in text.split("\n") if item] or [text] for text_part in text_parts: current_anchor = _append_paragraph_after( current_anchor, text_part, first_body_style, paragraph_template, ) return current_anchor def _replace_section_content( document: DocumentObject, anchor_title: str, target_title: str, content: dict, write_mode: str, after_element=None, ): heading = _find_heading_paragraph(document, anchor_title, after_element) if heading is None: return after_element _set_paragraph_text(heading, target_title, heading.style.name if heading.style is not None else None, heading) first_body_style, paragraph_template, table_template, blocks_to_remove = _collect_section_templates(heading) if write_mode == "replace_heading_only": return heading._element if write_mode == "replace_section": for block in blocks_to_remove: _delete_block(block) _insert_content_after( heading, content, first_body_style, paragraph_template, table_template, ) return heading._element def _group_logs(logs: list[dict]) -> list[list[dict]]: groups: list[list[dict]] = [] for item in logs: anchor_title = item.get("anchor_title") or item.get("title") or "" if not groups: groups.append([item]) continue last_group = groups[-1] last_anchor = last_group[0].get("anchor_title") or last_group[0].get("title") or "" if anchor_title == last_anchor: last_group.append(item) else: groups.append([item]) return groups def _replace_section_group( document: DocumentObject, items: list[dict], after_element=None, ): first_item = items[0] anchor_title = first_item.get("anchor_title") or first_item.get("title") or "" target_title = first_item.get("title") or anchor_title heading = _find_heading_paragraph(document, anchor_title, after_element) if heading is None: return after_element _set_paragraph_text(heading, target_title, heading.style.name if heading.style is not None else None, heading) first_body_style, paragraph_template, table_template, blocks_to_remove = _collect_section_templates(heading) if len(items) == 1 and first_item.get("write_mode") == "replace_heading_only": return heading._element preserve_existing = len(items) == 1 and first_item.get("write_mode") == "append_after_heading" if not preserve_existing: for block in blocks_to_remove: _delete_block(block) current_anchor = heading for item in items: current_anchor = _insert_content_after( current_anchor, item.get("content") or {"content": []}, first_body_style, paragraph_template, table_template, ) return heading._element def export_document_bytes(template_bytes: bytes, logs: list[dict]) -> bytes: document = Document(BytesIO(template_bytes)) referenced_anchors: set[str] = set() for item in logs: for key in ("anchor_title", "title"): val = (item.get(key) or "").strip() if val: referenced_anchors.add(val) print(f"[EXPORT] logs count={len(logs)}, anchor_titles={[(l.get('anchor_title'), l.get('title')) for l in logs]}") last_heading_element = None for group in _group_logs(logs): last_heading_element = _replace_section_group(document, group, last_heading_element) _remove_unreferenced_headings(document, referenced_anchors) output = BytesIO() document.save(output) return output.getvalue() def export_document_bytes_from_blocks(template_bytes: bytes, blocks: list[dict], logs: list[dict]) -> bytes: """基于块结构的文档导出,支持块排序。 blocks: TemplateBlock 序列化列表,每个包含 anchor_start_index、edit_mode、sort_index 等 logs: GenerationLog 序列化列表,每个包含 block_id、content、status """ document = Document(BytesIO(template_bytes)) # 构建 block_id → log 映射 log_by_block_id: dict[int, dict] = {} for log in logs: block_id = log.get("block_id") if block_id is not None: log_by_block_id[block_id] = log # 构建 block 映射 block_by_start_index: dict[int, dict] = {} for block in blocks: start_idx = block.get("anchor_start_index") if start_idx is not None: block_by_start_index[start_idx] = block # 按 sort_index 排序的块列表 sorted_blocks = sorted( [b for b in blocks if b.get("anchor_start_index") is not None], key=lambda b: b.get("sort_index", 0), ) # 枚举所有文档元素 all_items: list = list(_iter_block_items(document)) # 先对所有元素进行内容替换(在重排之前) for index, item in enumerate(all_items): block = block_by_start_index.get(index) if block is None: continue if block.get("edit_mode") == "manual": continue log = log_by_block_id.get(block.get("id")) if log is None or log.get("status") != "success": continue content = log.get("content") or {"content": []} if isinstance(content, str): import json try: content = json.loads(content) except Exception: content = {"content": [{"type": "text", "text": content}]} if isinstance(item, Paragraph): text_parts: list[str] = [] for c in content.get("content", []): if c.get("type") == "text": text_parts.append(c.get("text", "")) new_text = "\n".join(text_parts) if new_text: _set_paragraph_text(item, new_text, template_paragraph=item) # ---- 重排文档元素 ---- # 收集所有元素的 XML element,按块归类 body = document.element.body # 建立 index → element 映射 index_to_element: dict[int, any] = {} for index, item in enumerate(all_items): index_to_element[index] = item._element # 记录哪些索引已被块覆盖 covered_indices: set[int] = set() for block in blocks: start_idx = block.get("anchor_start_index") end_idx = block.get("anchor_end_index", start_idx) if start_idx is not None: for idx in range(start_idx, (end_idx or start_idx) + 1): if idx < len(all_items): covered_indices.add(idx) # 未被任何块覆盖的元素(gap),保持原顺序 gap_elements: list = [] for index, item in enumerate(all_items): if index not in covered_indices: gap_elements.append(item._element) # 块元素:按 sort_index 排序后收集 block_elements: list = [] for block in sorted_blocks: start_idx = block.get("anchor_start_index") end_idx = block.get("anchor_end_index", start_idx) if start_idx is not None: for idx in range(start_idx, (end_idx or start_idx) + 1): if idx < len(all_items) and idx in covered_indices: block_elements.append(index_to_element[idx]) # 移除所有元素 for element in list(body): body.remove(element) # 按新顺序重新添加:块元素(按 sort_index)→ 间隙元素(保持原顺序) for element in block_elements: body.append(element) for element in gap_elements: body.append(element) output = BytesIO() document.save(output) return output.getvalue()