#!/usr/bin/env python3
"""Detect unsafe structural/style drift between an original and edited DOCX."""

from __future__ import annotations

import argparse
import hashlib
import json
import sys
import zipfile
from dataclasses import asdict, dataclass
from itertools import zip_longest
from pathlib import Path
from xml.etree import ElementTree as ET


W = "{http://schemas.openxmlformats.org/wordprocessingml/2006/main}"


@dataclass(frozen=True)
class ParagraphRecord:
    index: int
    text: str
    style: str
    num_id: str
    ilvl: str
    outline_level: str


def read_part(package: zipfile.ZipFile, name: str) -> bytes | None:
    try:
        return package.read(name)
    except KeyError:
        return None


def sha256(data: bytes | None) -> str | None:
    return hashlib.sha256(data).hexdigest() if data is not None else None


def value(element: ET.Element | None) -> str:
    return element.get(f"{W}val", "") if element is not None else ""


def paragraph_records(document_xml: bytes) -> list[ParagraphRecord]:
    root = ET.fromstring(document_xml)
    records: list[ParagraphRecord] = []
    for index, paragraph in enumerate(root.iter(f"{W}p")):
        ppr = paragraph.find(f"{W}pPr")
        style = value(ppr.find(f"{W}pStyle") if ppr is not None else None)
        num_pr = ppr.find(f"{W}numPr") if ppr is not None else None
        num_id = value(num_pr.find(f"{W}numId") if num_pr is not None else None)
        ilvl = value(num_pr.find(f"{W}ilvl") if num_pr is not None else None)
        outline = value(ppr.find(f"{W}outlineLvl") if ppr is not None else None)
        text = "".join(node.text or "" for node in paragraph.iter(f"{W}t"))
        records.append(ParagraphRecord(index, text, style, num_id, ilvl, outline))
    return records


def document_counts(document_xml: bytes) -> dict[str, int]:
    root = ET.fromstring(document_xml)
    tags = {
        "paragraphs": "p",
        "tables": "tbl",
        "drawings": "drawing",
        "bookmarks": "bookmarkStart",
        "fields": "fldChar",
        "content_controls": "sdt",
        "sections": "sectPr",
        "page_breaks": "br",
    }
    return {name: sum(1 for _ in root.iter(f"{W}{tag}")) for name, tag in tags.items()}


def package_signature(path: Path) -> dict[str, object]:
    with zipfile.ZipFile(path) as package:
        document_xml = read_part(package, "word/document.xml")
        if document_xml is None:
            raise ValueError(f"{path} has no word/document.xml")
        guarded_parts = {}
        for name in package.namelist():
            if name in {
                "word/styles.xml",
                "word/numbering.xml",
                "word/settings.xml",
                "word/fontTable.xml",
                "word/theme/theme1.xml",
            } or name.startswith(("word/header", "word/footer", "word/media/")):
                guarded_parts[name] = sha256(read_part(package, name))
        return {
            "paragraphs": paragraph_records(document_xml),
            "counts": document_counts(document_xml),
            "guarded_parts": guarded_parts,
        }


def looks_like_heading(record: ParagraphRecord) -> bool:
    style = record.style.lower()
    return bool(record.outline_level or style.startswith("heading") or style.startswith("标题"))


def audit(original: Path, edited: Path, allowed_parts: set[str], allow_structure_change: bool) -> dict[str, object]:
    source = package_signature(original)
    target = package_signature(edited)
    failures: list[dict[str, object]] = []
    warnings: list[dict[str, object]] = []
    changed: list[dict[str, object]] = []

    if source["counts"] != target["counts"] and not allow_structure_change:
        failures.append({"kind": "document_structure_changed", "source": source["counts"], "edited": target["counts"]})

    source_parts = source["guarded_parts"]
    target_parts = target["guarded_parts"]
    for name in sorted(set(source_parts) | set(target_parts)):
        if source_parts.get(name) != target_parts.get(name) and name not in allowed_parts:
            failures.append({"kind": "guarded_package_part_changed", "part": name})

    for old, new in zip_longest(source["paragraphs"], target["paragraphs"]):
        if old is None or new is None:
            continue
        if old.text == new.text and (old.style, old.num_id, old.ilvl, old.outline_level) == (
            new.style,
            new.num_id,
            new.ilvl,
            new.outline_level,
        ):
            continue
        item = {"index": old.index, "source": asdict(old), "edited": asdict(new)}
        changed.append(item)
        if (old.style, old.num_id, old.ilvl, old.outline_level) != (
            new.style,
            new.num_id,
            new.ilvl,
            new.outline_level,
        ):
            failures.append({"kind": "paragraph_format_signature_changed", **item})
        if (
            old.style
            and len(old.text.strip()) <= 30
            and len(new.text.strip()) >= 50
            and any(mark in new.text for mark in ("。", "；"))
        ):
            failures.append({"kind": "body_text_replaced_short_styled_paragraph", **item})
        if looks_like_heading(new) and (len(new.text) > 80 or "。" in new.text or "；" in new.text):
            failures.append({"kind": "long_sentence_in_heading", **item})
        elif looks_like_heading(new) and new.text.startswith(("一是", "二是", "三是", "四是")):
            warnings.append({"kind": "enumerated_body_text_in_heading", **item})

    return {
        "original": str(original),
        "edited": str(edited),
        "result": "FAIL" if failures else "PASS",
        "failures": failures,
        "warnings": warnings,
        "changed_paragraphs": changed,
    }


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("original", type=Path)
    parser.add_argument("edited", type=Path)
    parser.add_argument("--allow-structure-change", action="store_true")
    parser.add_argument("--allow-part-change", action="append", default=[], metavar="DOCX_PART")
    parser.add_argument("--report", type=Path, help="Write the full JSON report to this path")
    args = parser.parse_args()

    try:
        report = audit(args.original, args.edited, set(args.allow_part_change), args.allow_structure_change)
    except (OSError, ValueError, zipfile.BadZipFile, ET.ParseError) as exc:
        print(f"audit error: {exc}", file=sys.stderr)
        return 2

    payload = json.dumps(report, ensure_ascii=False, indent=2)
    if args.report:
        args.report.parent.mkdir(parents=True, exist_ok=True)
        args.report.write_text(payload + "\n", encoding="utf-8")
    print(f"result={report['result']} failures={len(report['failures'])} warnings={len(report['warnings'])} changed_paragraphs={len(report['changed_paragraphs'])}")
    for failure in report["failures"][:20]:
        suffix = failure.get("part", failure.get("index", ""))
        print(f"FAIL {failure['kind']} {suffix}")
    if len(report["failures"]) > 20:
        print(f"... {len(report['failures']) - 20} more failures; use --report for details")
    return 1 if report["failures"] else 0


if __name__ == "__main__":
    raise SystemExit(main())
