#!/usr/bin/env python3
"""
Codebase Mapper -- Lightweight architecture visualization for AI agents.

Scans a project directory for source files (JS/TS/Python/Rust), extracts
imports/exports, and produces a dependency graph + summary that agents
can read to understand the codebase structure without reading every file.

Output:
  - .codebase-map/MAP.md       -- Human/agent-readable architecture summary
  - .codebase-map/graph.json   -- Machine-readable dependency graph

Usage:
    codebase-map                           # scan current directory
    codebase-map /path/to/project          # scan specific directory
    codebase-map --format json             # output JSON instead of markdown
    codebase-map --depth 3                 # limit directory depth
"""

import sys
import os
import re
import json
from pathlib import Path
from collections import defaultdict

# --- Configuration ---
SUPPORTED_EXTENSIONS = {
    '.ts', '.tsx', '.js', '.jsx', '.mjs',  # JavaScript/TypeScript
    '.py',                                   # Python
    '.rs',                                   # Rust
    '.css', '.scss',                         # Styles
}

SKIP_DIRS = {
    'node_modules', '.git', '.next', 'dist', 'build', '.cache',
    '__pycache__', '.pytest_cache', 'target', '.lab', '.session',
    '.vault-search', '.obsidian', '.claude', 'deprecated',
    'codebase-snapshots', '_github_imports',
}

SKIP_FILES = {
    'package-lock.json', 'yarn.lock', 'pnpm-lock.yaml',
}


# ============================================================
# File Scanning
# ============================================================

def scan_files(root, max_depth=None):
    """Find all source files in the project."""
    root = Path(root).resolve()
    files = []

    for path in root.rglob('*'):
        if max_depth is not None:
            rel = path.relative_to(root)
            if len(rel.parts) > max_depth:
                continue

        if any(skip in path.parts for skip in SKIP_DIRS):
            continue
        if path.name in SKIP_FILES:
            continue
        if path.suffix in SUPPORTED_EXTENSIONS and path.is_file():
            files.append(path)

    return sorted(files)


# ============================================================
# Import Extraction
# ============================================================

def extract_imports_ts(content):
    """Extract imports from TypeScript/JavaScript files."""
    imports = []

    # ES6 imports: import X from 'Y', import { X } from 'Y', import 'Y'
    for match in re.finditer(r'''import\s+(?:(?:[\w*{}\s,]+)\s+from\s+)?['"]([^'"]+)['"]''', content):
        imports.append(match.group(1))

    # Dynamic imports: import('Y'), require('Y')
    for match in re.finditer(r'''(?:import|require)\s*\(\s*['"]([^'"]+)['"]\s*\)''', content):
        imports.append(match.group(1))

    return imports


def extract_imports_python(content):
    """Extract imports from Python files."""
    imports = []

    # import X, from X import Y
    for match in re.finditer(r'^(?:from\s+([\w.]+)\s+import|import\s+([\w.]+))', content, re.MULTILINE):
        mod = match.group(1) or match.group(2)
        imports.append(mod)

    return imports


def extract_imports_rust(content):
    """Extract imports from Rust files."""
    imports = []

    # use X::Y
    for match in re.finditer(r'^use\s+([\w:]+)', content, re.MULTILINE):
        imports.append(match.group(1))

    return imports


def extract_exports_ts(content):
    """Extract exports from TypeScript/JavaScript files."""
    exports = []

    # export function/const/class/type/interface/enum
    for match in re.finditer(r'export\s+(?:default\s+)?(?:function|const|let|var|class|type|interface|enum)\s+(\w+)', content):
        exports.append(match.group(1))

    # export { X, Y }
    for match in re.finditer(r'export\s*\{([^}]+)\}', content):
        names = [n.strip().split(' as ')[0].strip() for n in match.group(1).split(',')]
        exports.extend(n for n in names if n)

    return exports


def extract_exports_python(content):
    """Extract exports from Python files (classes, functions, __all__)."""
    exports = []

    # __all__ = [...]
    all_match = re.search(r'__all__\s*=\s*\[([^\]]+)\]', content)
    if all_match:
        names = re.findall(r"'(\w+)'|\"(\w+)\"", all_match.group(1))
        exports.extend(n[0] or n[1] for n in names)
    else:
        # Top-level def and class
        for match in re.finditer(r'^(?:def|class)\s+(\w+)', content, re.MULTILINE):
            name = match.group(1)
            if not name.startswith('_'):
                exports.append(name)

    return exports


def analyze_file(filepath, root):
    """Analyze a single file for imports, exports, and metadata."""
    rel_path = str(filepath.relative_to(root))
    suffix = filepath.suffix

    try:
        content = filepath.read_text(encoding='utf-8', errors='ignore')
    except Exception:
        return None

    line_count = content.count('\n') + 1

    # Extract imports and exports based on language
    if suffix in ('.ts', '.tsx', '.js', '.jsx', '.mjs'):
        imports = extract_imports_ts(content)
        exports = extract_exports_ts(content)
        language = 'typescript' if suffix in ('.ts', '.tsx') else 'javascript'
    elif suffix == '.py':
        imports = extract_imports_python(content)
        exports = extract_exports_python(content)
        language = 'python'
    elif suffix == '.rs':
        imports = extract_imports_rust(content)
        exports = []  # Rust exports are more complex (pub items)
        language = 'rust'
    elif suffix in ('.css', '.scss'):
        imports = []
        exports = []
        language = 'css'
    else:
        imports = []
        exports = []
        language = 'unknown'

    # Classify imports as local vs external
    local_imports = []
    external_imports = []
    for imp in imports:
        if imp.startswith('.') or imp.startswith('/') or imp.startswith('~'):
            local_imports.append(imp)
        else:
            external_imports.append(imp)

    return {
        'path': rel_path,
        'language': language,
        'lines': line_count,
        'imports': imports,
        'local_imports': local_imports,
        'external_imports': external_imports,
        'exports': exports,
    }


# ============================================================
# Graph Building
# ============================================================

def resolve_import(importing_file, import_path, all_files):
    """Try to resolve a relative import to an actual file path."""
    if not import_path.startswith('.'):
        return None

    base_dir = Path(importing_file).parent
    resolved = (base_dir / import_path).resolve()

    # Try with common extensions
    candidates = [
        str(resolved),
        str(resolved) + '.ts',
        str(resolved) + '.tsx',
        str(resolved) + '.js',
        str(resolved) + '.jsx',
        str(resolved) + '.py',
        str(resolved / 'index.ts'),
        str(resolved / 'index.tsx'),
        str(resolved / 'index.js'),
        str(resolved / '__init__.py'),
    ]

    for candidate in candidates:
        for f in all_files:
            if str(f).endswith(candidate.split('/')[-1]) and candidate.replace('./', '') in str(f):
                return str(f)

    return None


def build_dependency_graph(file_analyses, root):
    """Build a dependency graph from file analyses."""
    all_paths = {a['path'] for a in file_analyses}
    edges = []

    for analysis in file_analyses:
        for imp in analysis['local_imports']:
            resolved = resolve_import(analysis['path'], imp, all_paths)
            if resolved:
                edges.append({
                    'from': analysis['path'],
                    'to': resolved,
                    'import': imp,
                })

    return edges


# ============================================================
# Output Generation
# ============================================================

def generate_directory_tree(file_analyses, root):
    """Generate a compact directory tree."""
    dirs = defaultdict(list)
    for a in file_analyses:
        parts = Path(a['path']).parts
        dir_path = '/'.join(parts[:-1]) if len(parts) > 1 else '.'
        dirs[dir_path].append({
            'name': parts[-1],
            'lines': a['lines'],
            'exports': len(a['exports']),
            'language': a['language'],
        })
    return dict(dirs)


def generate_markdown(file_analyses, edges, root):
    """Generate the MAP.md file."""
    lines = []
    lines.append("# Codebase Map")
    lines.append(f"\n**Project:** {Path(root).name}")
    lines.append(f"**Generated:** {__import__('datetime').datetime.now().strftime('%Y-%m-%d %H:%M')}")
    lines.append(f"**Files:** {len(file_analyses)}")
    total_lines = sum(a['lines'] for a in file_analyses)
    lines.append(f"**Total Lines:** {total_lines:,}")
    lines.append("")

    # Language breakdown
    lang_stats = defaultdict(lambda: {'files': 0, 'lines': 0})
    for a in file_analyses:
        lang_stats[a['language']]['files'] += 1
        lang_stats[a['language']]['lines'] += a['lines']

    lines.append("## Language Breakdown")
    lines.append("")
    lines.append("| Language | Files | Lines |")
    lines.append("|----------|-------|-------|")
    for lang, stats in sorted(lang_stats.items(), key=lambda x: -x[1]['lines']):
        lines.append(f"| {lang} | {stats['files']} | {stats['lines']:,} |")
    lines.append("")

    # Directory structure with file counts
    dir_tree = generate_directory_tree(file_analyses, root)
    lines.append("## Directory Structure")
    lines.append("")
    for dir_path in sorted(dir_tree.keys()):
        files = dir_tree[dir_path]
        total = sum(f['lines'] for f in files)
        lines.append(f"### `{dir_path}/` ({len(files)} files, {total:,} lines)")
        lines.append("")
        for f in sorted(files, key=lambda x: x['name']):
            export_note = f" -- exports: {', '.join(f.get('export_names', []))}" if f.get('export_names') else ""
            lines.append(f"- `{f['name']}` ({f['lines']} lines, {f['language']})")
        lines.append("")

    # External dependencies
    all_external = set()
    for a in file_analyses:
        all_external.update(a['external_imports'])

    if all_external:
        lines.append("## External Dependencies")
        lines.append("")
        # Group by prefix
        groups = defaultdict(list)
        for dep in sorted(all_external):
            prefix = dep.split('/')[0].lstrip('@')
            groups[prefix].append(dep)

        for prefix in sorted(groups.keys()):
            deps = groups[prefix]
            if len(deps) == 1:
                lines.append(f"- `{deps[0]}`")
            else:
                lines.append(f"- `{prefix}` ({len(deps)} imports)")
        lines.append("")

    # Key exports (public API surface)
    lines.append("## Key Exports")
    lines.append("")
    lines.append("| File | Exports |")
    lines.append("|------|---------|")
    for a in sorted(file_analyses, key=lambda x: x['path']):
        if a['exports']:
            exports_str = ', '.join(a['exports'][:5])
            if len(a['exports']) > 5:
                exports_str += f" (+{len(a['exports'])-5} more)"
            lines.append(f"| `{a['path']}` | {exports_str} |")
    lines.append("")

    # Internal dependencies (edges)
    if edges:
        lines.append("## Internal Dependencies")
        lines.append("")
        lines.append("```")
        # Group by source file
        by_source = defaultdict(list)
        for e in edges:
            by_source[e['from']].append(e['to'])
        for source in sorted(by_source.keys()):
            targets = by_source[source]
            for target in sorted(targets):
                lines.append(f"{source} --> {target}")
        lines.append("```")
        lines.append("")

    # Most connected files (highest dependency count)
    if edges:
        dep_count = defaultdict(int)
        for e in edges:
            dep_count[e['to']] += 1
        if dep_count:
            lines.append("## Most Imported Files")
            lines.append("")
            lines.append("| File | Imported By |")
            lines.append("|------|------------|")
            for path, count in sorted(dep_count.items(), key=lambda x: -x[1])[:10]:
                lines.append(f"| `{path}` | {count} files |")
            lines.append("")

    return '\n'.join(lines)


def generate_json(file_analyses, edges, root):
    """Generate the graph.json file."""
    return json.dumps({
        'project': Path(root).name,
        'files': file_analyses,
        'edges': edges,
        'stats': {
            'total_files': len(file_analyses),
            'total_lines': sum(a['lines'] for a in file_analyses),
        }
    }, indent=2)


# ============================================================
# CLI
# ============================================================

def main():
    args = sys.argv[1:]

    # Parse arguments
    root = '.'
    output_format = 'markdown'
    max_depth = None

    i = 0
    while i < len(args):
        if args[i] == '--format' and i + 1 < len(args):
            output_format = args[i + 1]
            i += 2
        elif args[i] == '--depth' and i + 1 < len(args):
            max_depth = int(args[i + 1])
            i += 2
        elif args[i] == '--help':
            print(__doc__)
            sys.exit(0)
        elif not args[i].startswith('--'):
            root = args[i]
            i += 1
        else:
            i += 1

    root = Path(root).resolve()
    if not root.exists():
        print(f"ERROR: Directory not found: {root}", file=sys.stderr)
        sys.exit(1)

    # Scan and analyze
    print(f"Scanning {root}...", file=sys.stderr)
    files = scan_files(root, max_depth)
    print(f"Found {len(files)} source files", file=sys.stderr)

    analyses = []
    for f in files:
        result = analyze_file(f, root)
        if result:
            analyses.append(result)

    edges = build_dependency_graph(analyses, root)
    print(f"Found {len(edges)} internal dependencies", file=sys.stderr)

    # Generate output
    output_dir = root / '.codebase-map'
    output_dir.mkdir(exist_ok=True)

    if output_format == 'json':
        output = generate_json(analyses, edges, root)
        print(output)
    else:
        md = generate_markdown(analyses, edges, root)
        md_path = output_dir / 'MAP.md'
        md_path.write_text(md)
        print(f"Map written to {md_path}", file=sys.stderr)

        graph = generate_json(analyses, edges, root)
        graph_path = output_dir / 'graph.json'
        graph_path.write_text(graph)
        print(f"Graph written to {graph_path}", file=sys.stderr)

        # Also print the markdown
        print(md)


if __name__ == "__main__":
    main()
