#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
make_mini_pkg.py

entry から import を静的解析し、指定パッケージ(pkg_name)配下に解決できた .py を再帰収集して
out 以下へ pkg_root からの相対パスを保ってコピーします（mini-packaging 用）。

追加機能:
(1) --report-unused : pkg_root 配下の .py のうち、選別(selected)されなかったものを unused_files.txt に出力
(2) --dot-out FILE  : 依存関係を Graphviz dot で出力（A -> B は A が B を import）

Defaults (requested):
- pkg_name default: "tklib"
- pkg_root default:
    1) join(env['tkProg_Root'], 'tklib', 'python', 'tklib') が存在すればそれ
    2) 'd:/git/tkProg/tklib/python/tklib' が存在すればそれ
    3) それでもダメなら entry から上方向に pkg_name フォルダを探索

Limitations:
- 動的 import(importlib 等) は検出できません
- from X import name の name が属性の場合は完全には追えません（X/__init__.py を保険で含めます）



1) デフォルトでtklibを検索
python make_mini_tklib.py --entry Ne-T_fit.py --out tklib_sub --verbose

2) tklib の場所を明示する場合
python make_mini_tklib.py --entry Ne-T_fit.py --pkg-root D:\path\to\tklib --out tklib_sub --verbose

3) パッケージ名を指定 の場所を明示する場合
python make_mini_tklib.py --entry optimize_mup.py --pkg-name tklib --pkg-root D:\git\tkProg\tklib\python\tklib --out tklib_sub --verbose

4) まずは何がコピーされるかだけ確認
python make_mini_tklib.py --entry optimize_mup.py --pkg-name tklib --pkg-root D:\git\tkProg\tklib\python\tklib --out tklib_sub --dry-run --verbose

5) ifでimportするモジュールをコピーしない
--ignore-if-imports

6) 未使用ライブラリを報告
--report-unused

7) 依存関係を可視化
--dot-out deps.dot
"""

from __future__ import annotations

import argparse
import ast
import os
import shutil
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Dict, List, Optional, Set, Tuple


# -----------------------------
# Path helpers
# -----------------------------
def norm(p: Path) -> Path:
    return p.expanduser().resolve()


def is_within(child: Path, parent: Path) -> bool:
    try:
        norm(child).relative_to(norm(parent))
        return True
    except Exception:
        return False


def find_pkg_root_upward(entry: Path, pkg_name: str) -> Optional[Path]:
    cur = norm(entry).parent
    for _ in range(50):
        cand = cur / pkg_name
        if cand.exists() and cand.is_dir():
            return norm(cand)
        if cur.parent == cur:
            break
        cur = cur.parent
    return None


def default_pkg_root_from_env_or_fallback() -> Optional[Path]:
    tkprog_root = os.environ.get("tkProg_Root")
    if tkprog_root:
        p1 = Path(os.path.join(tkprog_root, "tklib", "python", "tklib"))
        if p1.exists() and p1.is_dir():
            return norm(p1)

    p2 = Path("d:/git/tkProg/tklib/python/tklib")
    if p2.exists() and p2.is_dir():
        return norm(p2)

    return None


def iter_py_files_under(root: Path) -> List[Path]:
    return [p for p in norm(root).rglob("*.py") if p.is_file()]


# -----------------------------
# Import extraction
# -----------------------------
@dataclass(frozen=True)
class ImportRef:
    module: Optional[str]
    names: Optional[Tuple[str, ...]]  # for from-import
    level: int  # 0 absolute, >0 relative


def extract_imports(pyfile: Path, ignore_if_imports: bool = False) -> List[ImportRef]:
    """
    Extract imports from a python file.
    If ignore_if_imports=True, imports inside any `if ...:` block are ignored.
    """
    try:
        src = pyfile.read_text(encoding="utf-8")
    except UnicodeDecodeError:
        src = pyfile.read_text(encoding="utf-8", errors="replace")

    try:
        tree = ast.parse(src, filename=str(pyfile))
    except SyntaxError:
        return []

    out: List[ImportRef] = []

    class ImportCollector(ast.NodeVisitor):
        def __init__(self) -> None:
            self.if_depth = 0

        def _should_collect(self) -> bool:
            return not (ignore_if_imports and self.if_depth > 0)

        def visit_If(self, node: ast.If) -> None:
            self.if_depth += 1
            for stmt in node.body:
                self.visit(stmt)
            for stmt in node.orelse:
                self.visit(stmt)
            self.if_depth -= 1
            # do not visit node.test (imports won't be there)

        def visit_Import(self, node: ast.Import) -> None:
            if self._should_collect():
                for alias in node.names:
                    if alias.name:
                        out.append(ImportRef(module=alias.name, names=None, level=0))

        def visit_ImportFrom(self, node: ast.ImportFrom) -> None:
            if self._should_collect():
                mod = node.module
                lvl = int(node.level or 0)
                names = tuple(a.name for a in node.names if a.name)
                out.append(ImportRef(module=mod, names=names if names else None, level=lvl))

        def generic_visit(self, node: ast.AST) -> None:
            super().generic_visit(node)

    ImportCollector().visit(tree)
    return out


# -----------------------------
# Module resolution
# -----------------------------
def module_to_candidate_paths_under_root(pkg_root: Path, pkg_name: str, module: str) -> List[Path]:
    """
    module: "pkg.foo.bar" or "foo.bar"
    If module starts with pkg_name, strip it because pkg_root corresponds to pkg_name package dir.
    """
    parts = module.split(".")
    if parts and parts[0] == pkg_name:
        parts = parts[1:]

    cand1 = pkg_root.joinpath(*parts).with_suffix(".py")
    cand2 = pkg_root.joinpath(*parts) / "__init__.py"
    return [cand1, cand2]


def resolve_from_import_targets(
    pkg_root: Path,
    pkg_name: str,
    base_module: str,
    names: Tuple[str, ...],
) -> List[Path]:
    out: List[Path] = []

    # base module itself
    out.extend([p for p in module_to_candidate_paths_under_root(pkg_root, pkg_name, base_module) if p.exists()])

    # base_module.name as submodule/package
    for nm in names:
        out.extend(
            [p for p in module_to_candidate_paths_under_root(pkg_root, pkg_name, f"{base_module}.{nm}") if p.exists()]
        )

    return out


def resolve_import_ref_to_files(
    ref: ImportRef,
    pkg_root: Path,
    pkg_name: str,
    current_file: Path,
) -> List[Path]:
    files: List[Path] = []

    # Relative imports (filesystem-based)
    if ref.level and ref.level > 0:
        base_dir = current_file.parent
        for _ in range(ref.level - 1):
            base_dir = base_dir.parent

        if is_within(base_dir, pkg_root):
            target_dir = base_dir
            if ref.module:
                target_dir = target_dir.joinpath(*ref.module.split("."))

            if ref.names:
                for nm in ref.names:
                    cand_file = target_dir / f"{nm}.py"
                    cand_pkg = target_dir / nm / "__init__.py"
                    if cand_file.exists():
                        files.append(cand_file)
                    if cand_pkg.exists():
                        files.append(cand_pkg)

                initp = target_dir / "__init__.py"
                if initp.exists():
                    files.append(initp)

        return files

    # Absolute imports
    if not ref.module:
        return files

    # import X
    if ref.names is None:
        for p in module_to_candidate_paths_under_root(pkg_root, pkg_name, ref.module):
            if p.exists():
                files.append(p)
        return files

    # from X import a,b
    files.extend(resolve_from_import_targets(pkg_root, pkg_name, ref.module, ref.names))
    return files


# -----------------------------
# Copy logic
# -----------------------------
def rel_under_root(pkg_root: Path, file_path: Path) -> Path:
    return norm(file_path).relative_to(norm(pkg_root))


def ensure_parent_init_files(pkg_root: Path, file_path: Path, selected: Set[Path]) -> None:
    rel = rel_under_root(pkg_root, file_path)
    parts = list(rel.parts)
    if len(parts) <= 1:
        return

    for i in range(1, len(parts)):
        d = pkg_root.joinpath(*parts[:i])
        initp = d / "__init__.py"
        if initp.exists():
            selected.add(norm(initp))


def copy_selected_files(pkg_root: Path, out_root: Path, selected: Set[Path], dry_run: bool = False) -> None:
    out_root.mkdir(parents=True, exist_ok=True)

    for src in sorted(selected):
        if not is_within(src, pkg_root):
            continue
        rel = rel_under_root(pkg_root, src)
        dst = out_root / rel
        if dry_run:
            print(f"[DRY] {src} -> {dst}")
            continue
        dst.parent.mkdir(parents=True, exist_ok=True)
        shutil.copy2(src, dst)


# -----------------------------
# Reports
# -----------------------------
def write_unused_report(pkg_root: Path, selected: Set[Path], out_dir: Path, filename: str = "unused_files.txt") -> Path:
    """
    Write unused python files under pkg_root that are not in selected.
    Output paths are relative to pkg_root.
    """
    pkg_root = norm(pkg_root)
    out_dir = norm(out_dir)
    out_dir.mkdir(parents=True, exist_ok=True)

    all_py = {norm(p) for p in iter_py_files_under(pkg_root)}
    unused = sorted(all_py - {norm(p) for p in selected})

    out_path = out_dir / filename
    lines: List[str] = []
    lines.append(f"# pkg_root = {pkg_root}")
    lines.append(f"# total_py = {len(all_py)}")
    lines.append(f"# selected = {len(selected)}")
    lines.append(f"# unused   = {len(unused)}")
    lines.append("")
    for p in unused:
        try:
            rel = p.relative_to(pkg_root)
        except Exception:
            rel = p
        lines.append(str(rel).replace("\\", "/"))

    out_path.write_text("\n".join(lines), encoding="utf-8")
    return out_path


def write_dot_graph(
    pkg_root: Path,
    entry: Path,
    edges: Dict[Path, Set[Path]],
    out_path: Path,
) -> Path:
    """
    Write a Graphviz dot file.
    Nodes are module file paths relative to pkg_root when possible; otherwise absolute.
    """
    pkg_root = norm(pkg_root)
    entry = norm(entry)
    out_path = Path(out_path).expanduser().resolve()
    out_path.parent.mkdir(parents=True, exist_ok=True)

    def label(p: Path) -> str:
        p = norm(p)
        if is_within(p, pkg_root):
            s = str(p.relative_to(pkg_root)).replace("\\", "/")
        else:
            s = str(p).replace("\\", "/")
        return s

    lines: List[str] = []
    lines.append("digraph deps {")
    lines.append('  graph [rankdir="LR"];')
    lines.append('  node  [shape="box", fontsize=10];')
    lines.append("")

    # highlight entry
    entry_label = label(entry)
    lines.append(f'  "{entry_label}" [shape="box", style="rounded,bold"];')

    # nodes + edges
    for src, dsts in sorted(edges.items(), key=lambda kv: label(kv[0])):
        src_lab = label(src)
        for dst in sorted(dsts, key=label):
            dst_lab = label(dst)
            lines.append(f'  "{src_lab}" -> "{dst_lab}";')

    lines.append("}")
    out_path.write_text("\n".join(lines), encoding="utf-8")
    return out_path


# -----------------------------
# Main crawl
# -----------------------------
def build_mini_pkg(
    entry: Path,
    pkg_root: Path,
    pkg_name: str,
    out_root: Path,
    dry_run: bool = False,
    verbose: bool = False,
    ignore_if_imports: bool = False,
) -> Tuple[Set[Path], Dict[Path, Set[Path]]]:
    """
    Returns:
      selected_files: files under pkg_root that are required
      edges: dependency edges among selected files (src -> dst), both Path
    """
    entry = norm(entry)
    pkg_root = norm(pkg_root)
    out_root = out_root.expanduser().resolve()

    to_scan: List[Path] = [entry]
    scanned: Set[Path] = set()
    selected: Set[Path] = set()
    edges: Dict[Path, Set[Path]] = {}

    while to_scan:
        cur = norm(to_scan.pop())
        if cur in scanned:
            continue
        scanned.add(cur)

        if verbose:
            print(f"[SCAN] {cur}")

        if cur.suffix.lower() != ".py":
            continue

        # parse imports
        refs = extract_imports(cur, ignore_if_imports=ignore_if_imports)

        for ref in refs:
            files = resolve_import_ref_to_files(ref, pkg_root, pkg_name, cur)
            if not files:
                continue

            # only consider dst within pkg_root
            dsts: List[Path] = []
            for f in files:
                f = norm(f)
                if is_within(f, pkg_root):
                    dsts.append(f)

            if not dsts:
                continue

            # if cur itself is within pkg_root, record edges
            if is_within(cur, pkg_root):
                edges.setdefault(cur, set()).update(dsts)

            # add selected and enqueue
            for f in dsts:
                if f not in selected:
                    selected.add(f)
                    ensure_parent_init_files(pkg_root, f, selected)
                    if f not in scanned:
                        to_scan.append(f)

    copy_selected_files(pkg_root, out_root, selected, dry_run=dry_run)
    return selected, edges


def main() -> int:
    ap = argparse.ArgumentParser(
        description="Generate a minimal subset of a package used by an entry script, preserving tree under output root."
    )
    ap.add_argument("--entry", required=True, help="Path to entry (parent) Python script, e.g. Ne-T_fit.py")
    ap.add_argument("--pkg-name", default="tklib", help='Target package name (default: "tklib")')
    ap.add_argument(
        "--pkg-root",
        default=None,
        help="Path to the package directory (the folder that contains __init__.py). "
             "If omitted, uses tkProg_Root-based default then fallback then upward search.",
    )
    ap.add_argument("--out", default="tklib_sub", help="Output folder to write subset tree (default: tklib_sub)")
    ap.add_argument("--dry-run", action="store_true", help="Print planned copies but do not write files")
    ap.add_argument("--verbose", action="store_true", help="Verbose scan logs")

    # minimum-oriented options
    ap.add_argument(
        "--ignore-if-imports",
        action="store_true",
        help="Ignore imports that appear inside any 'if' block (more minimal subset).",
    )

    # reports
    ap.add_argument(
        "--report-unused",
        action="store_true",
        help="Write unused_files.txt under --out (all .py under pkg_root that were NOT selected).",
    )
    ap.add_argument(
        "--dot-out",
        default=None,
        help="Write Graphviz dot file of dependency edges among selected files (e.g. deps.dot).",
    )

    args = ap.parse_args()

    entry = Path(args.entry).expanduser()
    if not entry.exists():
        print(f"ERROR: entry not found: {entry}", file=sys.stderr)
        return 2

    pkg_name = args.pkg_name

    # Resolve pkg_root with requested defaults
    pkg_root: Optional[Path]
    if args.pkg_root:
        cand = Path(args.pkg_root).expanduser()
        if not (cand.exists() and cand.is_dir()):
            print(f"ERROR: --pkg-root not found or not a directory: {cand}", file=sys.stderr)
            return 2
        pkg_root = norm(cand)
    else:
        pkg_root = default_pkg_root_from_env_or_fallback()
        if pkg_root is None:
            pkg_root = find_pkg_root_upward(entry, pkg_name)

    if pkg_root is None or not (pkg_root.exists() and pkg_root.is_dir()):
        print(
            "ERROR: Could not determine pkg_root.\n"
            "  Provide --pkg-root explicitly, or set tkProg_Root env, or place the package folder in parents of entry.",
            file=sys.stderr,
        )
        return 2

    out_root = Path(args.out).expanduser()

    selected, edges = build_mini_pkg(
        entry=entry,
        pkg_root=pkg_root,
        pkg_name=pkg_name,
        out_root=out_root,
        dry_run=args.dry_run,
        verbose=args.verbose,
        ignore_if_imports=args.ignore_if_imports,
    )

    print(f"pkg_name : {pkg_name}")
    print(f"pkg_root : {pkg_root}")
    print(f"entry    : {norm(entry)}")
    print(f"out      : {out_root.resolve()}")
    print(f"selected : {len(selected)} file(s)")

    # Reports
    if args.report_unused:
        unused_path = write_unused_report(pkg_root, selected, out_root)
        print(f"unused_report : {unused_path}")

    if args.dot_out:
        dot_path = write_dot_graph(pkg_root, entry, edges, Path(args.dot_out))
        print(f"dot_graph    : {dot_path}")
        print("  (render example) dot -Tpng deps.dot -o deps.png")

    return 0


if __name__ == "__main__":
    raise SystemExit(main())