#!/usr/bin/env python3
"""ELA Course Audit Runner.

Usage:
    python3 audit/run_audit.py              # audit all weeks
    python3 audit/run_audit.py Week_02      # audit specific week
    python3 audit/run_audit.py --quick      # structural checks only
    python3 audit/run_audit.py --report     # generate full report file
    python3 audit/run_audit.py --db-review  # list words needing human review

Exit codes:
    0  — All checks passed
    1  — Errors found (see output)
    2  — Words need human review (warnings only)
"""

import os
import re
import sys
from datetime import datetime
from collections import defaultdict
from typing import Dict, List, Tuple

# Add parent directory for imports
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from audit.word_db import (
    get_word as lookup,
    get_all,
    get_review as needs_review,
    count,
    get_multi_patterns,
)
from audit.validators import validate_pattern
from audit.scanner import scan_all_weeks, scan_week, WordClaim, WeekData


class AuditResult:
    """Result of auditing a single word claim."""
    PASS = "pass"
    FAIL = "fail"
    REVIEW = "review"
    WARNING = "warning"

    def __init__(self, claim: WordClaim, status: str, message: str):
        self.claim = claim
        self.status = status
        self.message = message

    def __repr__(self):
        icon = {"pass": "✓", "fail": "✗", "review": "⚠", "warning": "?"}[self.status]
        return f"{icon} {self.claim.week} {self.claim.word}: {self.message}"


class AuditReport:
    """Collects and formats audit results."""

    def __init__(self):
        self.results: List[AuditResult] = []
        self.weeks_scanned: List[int] = []
        self.total_claims = 0
        self.start_time = datetime.now()

    def add(self, claim: WordClaim, status: str, message: str):
        self.results.append(AuditResult(claim, status, message))

    def summary(self) -> dict:
        counts = defaultdict(int)
        for r in self.results:
            counts[r.status] += 1
        return dict(counts)

    def print_report(self, verbose: bool = False):
        """Print formatted audit report."""
        elapsed = (datetime.now() - self.start_time).total_seconds()
        summary = self.summary()

        print("=" * 60)
        print("ELA COURSE AUDIT REPORT")
        print(f"Generated: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
        print(f"Elapsed: {elapsed:.1f}s")
        print("=" * 60)
        print()

        # Overall summary
        print("SUMMARY")
        print("-" * 40)
        print(f"  Weeks scanned:     {len(self.weeks_scanned)}")
        print(f"  Total word claims: {len(self.results)}")
        print(f"  ✓ Passed:          {summary.get('pass', 0)}")
        print(f"  ✗ Errors:          {summary.get('fail', 0)}")
        print(f"  ⚠ Need review:     {summary.get('review', 0)}")
        print(f"  ? Warnings:        {summary.get('warning', 0)}")
        print()

        # Errors first
        fails = [r for r in self.results if r.status == AuditResult.FAIL]
        if fails:
            print("ERRORS (must fix):")
            print("-" * 40)
            for r in fails:
                print(f"  {r}")
            print()

        # Reviews
        reviews = [r for r in self.results if r.status == AuditResult.REVIEW]
        if reviews:
            print("NEEDS HUMAN REVIEW:")
            print("-" * 40)
            for r in reviews:
                print(f"  {r}")
            print()

        # Warnings
        warnings = [r for r in self.results if r.status == AuditResult.WARNING]
        if warnings and verbose:
            print("WARNINGS:")
            print("-" * 40)
            for r in warnings:
                print(f"  {r}")
            print()

        # Per-week breakdown
        if verbose:
            print("PER-WEEK BREAKDOWN:")
            print("-" * 40)
            by_week: Dict[str, List[AuditResult]] = defaultdict(list)
            for r in self.results:
                by_week[r.claim.week].append(r)

            for week in sorted(by_week.keys()):
                week_results = by_week[week]
                week_pass = sum(1 for r in week_results if r.status == AuditResult.PASS)
                week_fail = sum(1 for r in week_results if r.status == AuditResult.FAIL)
                week_review = sum(1 for r in week_results if r.status == AuditResult.REVIEW)
                status = "OK" if week_fail == 0 else f"FAIL({week_fail})"
                if week_review:
                    status += f" REVIEW({week_review})"
                print(f"  {week}: {len(week_results)} words — {status}")

        print()

        # Word database stats
        db_review = needs_review()
        if db_review or verbose:
            print("WORD DATABASE:")
            print("-" * 40)
            print(f"  Total verified words: {count()}")
            if db_review:
                print(f"  Needs review: {len(db_review)}")
                for entry in db_review:
                    print(f"    • {entry.word} — {entry.notes}")
            print()

        return fails

    def save_report(self, path: str):
        """Save report to file."""
        with open(path, 'w') as f:
            f.write(self._markdown_report())
        print(f"\nReport saved to: {path}")

    def _markdown_report(self) -> str:
        lines = [
            "# ELA Course Audit Report",
            f"Generated: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}",
            "",
            "## Summary",
            f"- Weeks scanned: {len(self.weeks_scanned)}",
            f"- Total claims: {len(self.results)}",
        ]
        summary = self.summary()
        lines.append(f"- ✓ Passed: {summary.get('pass', 0)}")
        lines.append(f"- ✗ Errors: {summary.get('fail', 0)}")
        lines.append(f"- ⚠ Review: {summary.get('review', 0)}")
        lines.append(f"- ? Warnings: {summary.get('warning', 0)}")
        lines.append("")

        # Errors
        fails = [r for r in self.results if r.status == AuditResult.FAIL]
        if fails:
            lines.append("## Errors (Must Fix)")
            for r in fails:
                lines.append(f"- **{r.claim.word}** ({r.claim.week}): {r.message}")
            lines.append("")

        # Reviews
        reviews = [r for r in self.results if r.status == AuditResult.REVIEW]
        if reviews:
            lines.append("## Needs Human Review")
            for r in reviews:
                lines.append(f"- **{r.claim.word}** ({r.claim.week}): {r.message}")
            lines.append("")

        # Per-week
        lines.append("## Per-Week Breakdown")
        by_week: Dict[str, List[AuditResult]] = defaultdict(list)
        for r in self.results:
            by_week[r.claim.week].append(r)
        for week in sorted(by_week.keys()):
            wr = by_week[week]
            wp = sum(1 for r in wr if r.status == AuditResult.PASS)
            wf = sum(1 for r in wr if r.status == AuditResult.FAIL)
            lines.append(f"- **{week}**: {len(wr)} words, {wp} pass, {wf} errors")

        return "\n".join(lines)


def audit_claim(claim: WordClaim, db_only: bool = False) -> AuditResult:
    """Audit a single word claim.

    Checks:
    1. Structural validation (pattern matches word structure)
    2. Database verification (word exists in verified database)
    3. Pattern consistency (claimed pattern matches database pattern)
    4. Multi-pattern support (words taught from different angles in different weeks)
    """
    word = claim.word
    claimed = claim.claimed_pattern

    # Check 1: Structural validation
    if not db_only:
        valid, msg = validate_pattern(word, claimed)
        if not valid:
            return AuditResult(claim, AuditResult.FAIL, msg)

    # Check 2: Database lookup
    db_entry = lookup(word)
    if not db_entry:
        # Word not in database — flag for review
        return AuditResult(claim, AuditResult.REVIEW,
                         f"not in verified database — add to word_db.py")

    # Normalize claimed pattern
    normalized = _normalize_pattern(claimed)

    # Check 3a: Multi-pattern words — allow multiple valid patterns
    multi_patterns = get_multi_patterns(word)
    if multi_patterns:
        if normalized in multi_patterns:
            return AuditResult(claim, AuditResult.PASS,
                              f"verified (multi-pattern): {normalized}")
        # Claimed pattern not in allowed list
        return AuditResult(claim, AuditResult.WARNING,
                         f"claimed '{claimed}' but allowed patterns are {multi_patterns}")

    # Check 3b: Standard pattern match
    if db_entry.pattern != normalized:
        return AuditResult(claim, AuditResult.WARNING,
                         f"claimed '{claimed}' but database has '{db_entry.pattern}'")

    return AuditResult(claim, AuditResult.PASS,
                      f"verified: {db_entry.pattern}")


def _normalize_pattern(claimed: str) -> str:
    """Normalize a claimed pattern to database pattern code.

    Maps various claim formats to the canonical pattern codes used in word_db.
    """
    claimed = claimed.lower().strip()

    # Silent E patterns
    silent_e_map = {
        "long a (silent e)": "long_a_a_e",
        "long e (silent e)": "long_e_e_e",
        "long i (silent e)": "long_i_i_e",
        "long o (silent e)": "long_o_o_e",
        "long u (silent e)": "long_u_u_e",
    }
    if claimed in silent_e_map:
        return silent_e_map[claimed]

    # Short vowels
    if claimed in ["short a", "short e", "short i", "short o", "short u"]:
        return claimed.replace(" ", "_")

    # Vowel teams
    team_map = {
        "long a (ai)": "long_a_ai",
        "long e (ee)": "long_e_ee",
        "long e (ea)": "long_e_ea",
        "long a (oa)": "long_o_oa",
    }
    if claimed in team_map:
        return team_map[claimed]

    # Digraphs
    digraph_map = {
        "sh sound (sh)": "digraph_sh",
        "ch sound (ch)": "digraph_ch",
        "th sound (th)": "digraph_th",
        "wh sound (wh)": "digraph_wh",
        "ph sound (ph)": "digraph_ph",
        "ck sound (ck)": "digraph_ck",
    }
    if claimed in digraph_map:
        return digraph_map[claimed]

    # Soft C/G
    if claimed.startswith("soft c"):
        suffix = claimed.split("(")[-1].rstrip(")") if "(" in claimed else ""
        if suffix in ["ce"]:
            return "soft_c_ce"
        elif suffix in ["ci"]:
            return "soft_c_ci"
        elif suffix in ["cy"]:
            return "soft_c_cy"
        return "soft_c_" + suffix
    if claimed.startswith("soft g"):
        suffix = claimed.split("(")[-1].rstrip(")") if "(" in claimed else ""
        if suffix in ["ge"]:
            return "soft_g_ge"
        elif suffix in ["gi"]:
            return "soft_g_gi"
        elif suffix in ["gy"]:
            return "soft_g_gy"
        return "soft_g_" + suffix

    # R-controlled
    r_map = {
        "ar = /ar/": "r_control_ar",
        "or = /or/": "r_control_or",
        "er = /er/": "r_control_er",
        "ir = /er/": "r_control_ir",
        "ur = /er/": "r_control_ur",
    }
    if claimed in r_map:
        return r_map[claimed]

    # Suffix patterns (add before double consonant)
    # "base + ing" → suffix_ing
    if re.match(r'^.+ \+ [a-z]*ing$', claimed):
        return "suffix_ing"
    # "base + s" → suffix_s (exclude es/est which have their own handlers)
    if re.match(r'^.+ \+ s$', claimed):
        return "suffix_s"
    if re.match(r'^.+ \+ [^es]s$', claimed):
        return "suffix_s"
    # "base + es" → suffix_es
    if re.match(r'^.+ \+ .?es$', claimed):
        return "suffix_es"
    # "base + ed" → suffix_ed
    if re.match(r'^.+ \+ .?ed$', claimed):
        return "suffix_ed"
    # "base + er" → suffix_er
    if re.match(r'^.+ \+ .?er$', claimed):
        return "suffix_er"
    # "base + est" → suffix_est
    if re.match(r'^.+ \+ .?est$', claimed):
        return "suffix_est"
    # "base + ly" → suffix_ly
    if re.match(r'^.+ \+ .?ly$', claimed):
        return "suffix_ly"
    # "base + fully" → suffix_ly
    if re.match(r'^.+ \+ fully$', claimed):
        return "suffix_ly"
    # "base + ful" → suffix_ful
    if re.match(r'^.+ \+ ful$', claimed):
        return "suffix_ful"
    # "base + less" → suffix_less
    if re.match(r'^.+ \+ .?less$', claimed):
        return "suffix_less"

    # IE rule: "ie - i before e" → ie_rule
    if "ie" in claimed and "i before e" in claimed:
        return "ie_rule"

    # Double consonant (add before irregular)
    if "tt" in claimed or "pp" in claimed or "mm" in claimed or "nn" in claimed:
        if " = " in claimed or "irregular" in claimed:
            return "double_consonant"

    # Irregular suffixes: "irregular - ittle" → double_consonant
    if "irregular" in claimed and ("tt" in claimed or "pp" in claimed or "mm" in claimed or "nn" in claimed):
        return "double_consonant"
    # "irregular - other" → digraph_th
    if "irregular" in claimed and ("other" in claimed or "ather" in claimed):
        return "digraph_th"
    # "irregular - ogeth" → digraph_th (together)
    if "irregular" in claimed and "ogeth" in claimed:
        return "digraph_th"
    # "irregular - ead" → short_e (for head/bread)
    if "irregular" in claimed and "ead" in claimed:
        return "short_e"

    # Generic irregular
    if "irregular" in claimed:
        return "irregular"

    # Can't normalize — return as-is for comparison
    return claimed


def run_audit(target: str = None, quick: bool = False, verbose: bool = False) -> AuditReport:
    """Run the full audit.

    Args:
        target: Specific week directory to audit, or None for all weeks
        quick: If True, only run structural checks (skip database)
        verbose: If True, show detailed output

    Returns:
        AuditReport with all results
    """
    base_dir = os.path.expanduser("~/Home_School/2nd_Grade/English_Language_Arts")
    report = AuditReport()

    if target:
        target_path = os.path.join(base_dir, target)
        if not os.path.isdir(target_path):
            print(f"Error: {target_path} not found")
            sys.exit(1)
        data = scan_week(target_path)
        if data:
            report.weeks_scanned.append(data.week_num)
            _audit_week_data(data, report, quick)
    else:
        weeks = scan_all_weeks(base_dir)
        for data in weeks:
            report.weeks_scanned.append(data.week_num)
            _audit_week_data(data, report, quick)

    return report


def _audit_week_data(data: WeekData, report: AuditReport, quick: bool):
    """Audit all claims in a week's data."""
    for claim in data.word_claims:
        result = audit_claim(claim, db_only=quick)
        report.add(claim, result.status, result.message)

    # Also check scanner errors
    for error in data.errors:
        report.add(
            WordClaim("?", "unknown", error, f"Week_{data.week_num:02d}"),
            AuditResult.FAIL,
            error
        )


class PreGenerateError(Exception):
    """Raised when pre-generation validation fails."""
    pass


def validate_before_generate(generator_path: str, strict: bool = True):
    """Pre-generation validation hook.

    Import this into any week generator and call at the top of main().
    Blocks generation if validation fails.

    Args:
        generator_path: Path to the generator file (e.g., __file__)
        strict: If True, fail on errors. If False, warn only.

    Raises:
        PreGenerateError: If validation errors found and strict=True.

    Usage:
        from audit.run_audit import validate_before_generate
        def main():
            validate_before_generate(__file__)  # ← blocks on errors
            # ... rest of generation
    """
    generator_path = os.path.abspath(generator_path)
    week_dir = os.path.dirname(generator_path)

    data = scan_week(week_dir)
    if not data:
        if strict:
            raise PreGenerateError(f"Could not scan generator: {generator_path}")
        return

    report = AuditReport()
    report.weeks_scanned.append(data.week_num)
    _audit_week_data(data, report, quick=False)

    summary = report.summary()
    errors = summary.get("fail", 0)
    reviews = summary.get("review", 0)
    warnings = summary.get("warning", 0)

    if errors == 0 and reviews == 0:
        print(f"  ✓ Audit passed ({len(data.word_claims)} words verified)")
        return

    # Build error message
    msg_lines = []
    msg_lines.append(f"  ✗ Audit FAILED for Week {data.week_num:02d} ({errors} errors, {reviews} need review)")

    fails = [r for r in report.results if r.status == AuditResult.FAIL]
    for r in fails:
        msg_lines.append(f"    ✗ {r.claim.word}: {r.message}")

    reviews_list = [r for r in report.results if r.status == AuditResult.REVIEW]
    for r in reviews_list:
        msg_lines.append(f"    ⚠ {r.claim.word}: {r.message}")

    msg = "\n".join(msg_lines)

    if strict:
        raise PreGenerateError(msg + "\n\nFix the issues above before generating.")
    else:
        print(msg)


def main():
    import argparse

    parser = argparse.ArgumentParser(description="ELA Course Audit System")
    parser.add_argument("week", nargs="?", default=None,
                       help="Specific week to audit (e.g., Week_02)")
    parser.add_argument("--quick", action="store_true",
                       help="Structural checks only, skip database")
    parser.add_argument("--verbose", "-v", action="store_true",
                       help="Show detailed output")
    parser.add_argument("--report", action="store_true",
                       help="Save report to file")
    parser.add_argument("--db-review", action="store_true",
                       help="List words needing human review in database")

    args = parser.parse_args()

    if args.db_review:
        review = needs_review()
        if not review:
            print("No words need review.")
            return
        print(f"\nWords needing human review ({len(review)}):")
        for entry in review:
            print(f"  • {entry.word:15} pattern={entry.pattern:20} notes: {entry.notes}")
        return

    report = run_audit(args.week, quick=args.quick, verbose=args.verbose)
    fails = report.print_report(verbose=args.verbose)

    if args.report:
        report_dir = os.path.expanduser("~/Home_School/2nd_Grade/English_Language_Arts/audit/reports")
        os.makedirs(report_dir, exist_ok=True)
        timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
        report.save_report(os.path.join(report_dir, f"audit_{timestamp}.md"))

    # Exit code
    summary = report.summary()
    if summary.get("fail", 0) > 0:
        sys.exit(1)
    elif summary.get("review", 0) > 0:
        sys.exit(2)
    else:
        sys.exit(0)


if __name__ == "__main__":
    main()