#!/usr/bin/env python3
"""
Thuban Python AST Helper

Reads Python source from stdin, parses it with the built-in `ast` module,
and prints a single-line JSON report to stdout describing:

  - unused_import       Imported names never referenced anywhere in the file
  - unused_variable     Local variables assigned in a function but never read
  - high_complexity     Functions whose McCabe cyclomatic complexity exceeds
                         COMPLEXITY_THRESHOLD
  - dead_code           Statements that are unreachable because they follow
                         an unconditional return/raise/break/continue in the
                         same block
  - bare_except         `except:` clauses with no exception type
  - mutable_default_arg Function parameters defaulting to a mutable literal
                         (list/dict/set) or list()/dict()/set() call

This script is invoked by packages/scanner/python-ast-analyzer.js as a
subprocess. It must never raise past main() - on any parse failure it
prints {"ok": false, "error": "..."} and exits 0, so the caller can treat
AST analysis as a best-effort supplement.

Two invocation modes are supported:

  Single-file mode (default, backward compatible): stdin is the raw
  Python source of one file.
    Output (stdout, single JSON object):
      {"ok": true, "issues": [{"type": str, "line": int, "name": str|null, "message": str}, ...]}
      {"ok": false, "error": str}

  Batch mode (`--batch` argv flag): starting a fresh Python interpreter
  per file is the dominant cost of AST analysis (interpreter startup
  alone is tens of milliseconds, dwarfing the actual parse/walk work per
  file), so for multi-file scans the caller instead sends ALL files in a
  single subprocess invocation. stdin is a JSON array of
  [{"path": str, "content": str}, ...]; stdout is a single JSON object:
      {"ok": true, "results": [{"path": str, "ok": bool, "issues": [...], "error": str|null}, ...]}
  Each entry's shape mirrors the single-file contract above so callers
  can reuse the same per-file result handling either way.
"""

import ast
import json
import sys

COMPLEXITY_THRESHOLD = 10

TERMINATOR_TYPES = (ast.Return, ast.Raise, ast.Continue, ast.Break)


# ─────────────────────────────────────────────────────────────────────────
#  Unused imports
# ─────────────────────────────────────────────────────────────────────────

class _UnusedImportChecker:
    def __init__(self, tree):
        self.tree = tree
        self.imports = []          # [{'name', 'line', 'source'}]
        self.used_names = set()
        self.dunder_all_names = set()

    def collect(self):
        for node in ast.walk(self.tree):
            if isinstance(node, ast.Import):
                for alias in node.names:
                    if alias.name == '*':
                        continue
                    bound = alias.asname or alias.name.split('.')[0]
                    self.imports.append({'name': bound, 'line': node.lineno, 'source': ''})
            elif isinstance(node, ast.ImportFrom):
                if node.module == '__future__':
                    continue
                for alias in node.names:
                    if alias.name == '*':
                        continue
                    bound = alias.asname or alias.name
                    self.imports.append({'name': bound, 'line': node.lineno, 'source': node.module or ''})
            elif isinstance(node, ast.Name) and isinstance(node.ctx, ast.Load):
                self.used_names.add(node.id)
            elif isinstance(node, ast.Assign):
                for target in node.targets:
                    if isinstance(target, ast.Name) and target.id == '__all__' and \
                            isinstance(node.value, (ast.List, ast.Tuple, ast.Set)):
                        for elt in node.value.elts:
                            if isinstance(elt, ast.Constant) and isinstance(elt.value, str):
                                self.dunder_all_names.add(elt.value)

    def unused(self):
        issues = []
        seen = set()
        for imp in self.imports:
            name = imp['name']
            if name.startswith('_'):
                continue
            if name in self.used_names or name in self.dunder_all_names:
                continue
            key = (name, imp['line'])
            if key in seen:
                continue
            seen.add(key)
            source_suffix = " from '%s'" % imp['source'] if imp['source'] else ''
            issues.append({
                'type': 'unused_import',
                'line': imp['line'],
                'name': name,
                'message': "Unused import: '%s'%s" % (name, source_suffix),
            })
        return issues


# ─────────────────────────────────────────────────────────────────────────
#  Per-function analysis: complexity, unused locals, mutable defaults
# ─────────────────────────────────────────────────────────────────────────

class _ComplexityVisitor(ast.NodeVisitor):
    """Counts McCabe decision points within a single function's own body,
    stopping at nested function/lambda boundaries (those are analyzed
    independently when ast.walk() reaches them on its own)."""

    def __init__(self):
        self.complexity = 1

    def visit_If(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_For(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_AsyncFor(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_While(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_IfExp(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_ExceptHandler(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_BoolOp(self, node):
        self.complexity += max(len(node.values) - 1, 0)
        self.generic_visit(node)

    def visit_comprehension(self, node):
        self.complexity += len(node.ifs)
        self.generic_visit(node)

    # Nested scopes get their own complexity computation - don't descend.
    def visit_FunctionDef(self, node):
        pass

    def visit_AsyncFunctionDef(self, node):
        pass

    def visit_Lambda(self, node):
        pass


def _compute_complexity(fn):
    visitor = _ComplexityVisitor()
    for stmt in fn.body:
        visitor.visit(stmt)
    return visitor.complexity


def _is_mutable_default(default):
    if isinstance(default, (ast.List, ast.Dict, ast.Set)):
        return True
    if isinstance(default, ast.Call) and isinstance(default.func, ast.Name):
        return default.func.id in ('list', 'dict', 'set')
    return False


def _mutable_defaults(fn):
    issues = []
    args = fn.args
    positional = list(getattr(args, 'posonlyargs', [])) + list(args.args)
    if args.defaults:
        for arg, default in zip(reversed(positional), reversed(args.defaults)):
            if _is_mutable_default(default):
                issues.append({
                    'type': 'mutable_default_arg',
                    'line': default.lineno,
                    'name': arg.arg,
                    'message': "Mutable default argument '%s' in function '%s' - shared across all calls" % (arg.arg, fn.name),
                })
    for arg, default in zip(args.kwonlyargs, args.kw_defaults or []):
        if default is not None and _is_mutable_default(default):
            issues.append({
                'type': 'mutable_default_arg',
                'line': default.lineno,
                'name': arg.arg,
                'message': "Mutable default argument '%s' in function '%s' - shared across all calls" % (arg.arg, fn.name),
            })
    return issues


class _LocalAssignCollector(ast.NodeVisitor):
    """Collects simple (single-Name-target) local assignments made directly
    within a function's own body - not inside nested function/lambda
    scopes, which own their own locals."""

    def __init__(self):
        self.assigned = {}   # name -> first assignment line
        self.declared_global_nonlocal = set()

    def visit_Assign(self, node):
        if len(node.targets) == 1 and isinstance(node.targets[0], ast.Name):
            name = node.targets[0].id
            self.assigned.setdefault(name, node.lineno)
        self.generic_visit(node)

    def visit_AnnAssign(self, node):
        if isinstance(node.target, ast.Name) and node.value is not None:
            self.assigned.setdefault(node.target.id, node.lineno)
        self.generic_visit(node)

    def visit_For(self, node):
        if isinstance(node.target, ast.Name):
            self.assigned.setdefault(node.target.id, node.lineno)
        self.generic_visit(node)

    def visit_With(self, node):
        for item in node.items:
            if item.optional_vars is not None and isinstance(item.optional_vars, ast.Name):
                self.assigned.setdefault(item.optional_vars.id, node.lineno)
        self.generic_visit(node)

    def visit_Global(self, node):
        self.declared_global_nonlocal.update(node.names)

    def visit_Nonlocal(self, node):
        self.declared_global_nonlocal.update(node.names)

    def visit_FunctionDef(self, node):
        pass

    def visit_AsyncFunctionDef(self, node):
        pass

    def visit_Lambda(self, node):
        pass


class _LoadNameCollector(ast.NodeVisitor):
    """Collects every Name read (Load context) anywhere within a subtree,
    INCLUDING nested functions/lambdas - a variable captured by a closure
    still counts as used."""

    def __init__(self):
        self.loaded = set()

    def visit_Name(self, node):
        if isinstance(node.ctx, ast.Load):
            self.loaded.add(node.id)
        self.generic_visit(node)


def _unused_locals(fn):
    assign_collector = _LocalAssignCollector()
    for stmt in fn.body:
        assign_collector.visit(stmt)

    load_collector = _LoadNameCollector()
    load_collector.visit(fn)

    issues = []
    for name, line in assign_collector.assigned.items():
        if name.startswith('_'):
            continue
        if name in assign_collector.declared_global_nonlocal:
            continue
        if name in load_collector.loaded:
            continue
        issues.append({
            'type': 'unused_variable',
            'line': line,
            'name': name,
            'message': "Variable '%s' is assigned but never used" % name,
        })
    return issues


# ─────────────────────────────────────────────────────────────────────────
#  Dead code (unreachable statements)
# ─────────────────────────────────────────────────────────────────────────

def _check_block(stmts, issues):
    for i, stmt in enumerate(stmts):
        if isinstance(stmt, TERMINATOR_TYPES):
            if i + 1 < len(stmts):
                nxt = stmts[i + 1]
                kind = type(stmt).__name__.lower()
                issues.append({
                    'type': 'dead_code',
                    'line': nxt.lineno,
                    'name': None,
                    'message': "Unreachable code after '%s' statement" % kind,
                })
            break  # only report the first unreachable run per block


class _DeadCodeVisitor(ast.NodeVisitor):
    def __init__(self):
        self.issues = []

    def visit_Module(self, node):
        _check_block(node.body, self.issues)
        self.generic_visit(node)

    def visit_FunctionDef(self, node):
        _check_block(node.body, self.issues)
        self.generic_visit(node)

    def visit_AsyncFunctionDef(self, node):
        _check_block(node.body, self.issues)
        self.generic_visit(node)

    def visit_If(self, node):
        _check_block(node.body, self.issues)
        _check_block(node.orelse, self.issues)
        self.generic_visit(node)

    def visit_For(self, node):
        _check_block(node.body, self.issues)
        _check_block(node.orelse, self.issues)
        self.generic_visit(node)

    def visit_AsyncFor(self, node):
        _check_block(node.body, self.issues)
        _check_block(node.orelse, self.issues)
        self.generic_visit(node)

    def visit_While(self, node):
        _check_block(node.body, self.issues)
        _check_block(node.orelse, self.issues)
        self.generic_visit(node)

    def visit_Try(self, node):
        _check_block(node.body, self.issues)
        for handler in node.handlers:
            _check_block(handler.body, self.issues)
        _check_block(node.orelse, self.issues)
        _check_block(node.finalbody, self.issues)
        self.generic_visit(node)

    def visit_With(self, node):
        _check_block(node.body, self.issues)
        self.generic_visit(node)

    def visit_AsyncWith(self, node):
        _check_block(node.body, self.issues)
        self.generic_visit(node)


# ─────────────────────────────────────────────────────────────────────────
#  Suspicious patterns: bare except
# ─────────────────────────────────────────────────────────────────────────

def _find_bare_except(tree):
    issues = []
    for node in ast.walk(tree):
        if isinstance(node, ast.ExceptHandler) and node.type is None:
            issues.append({
                'type': 'bare_except',
                'line': node.lineno,
                'name': None,
                'message': "Bare 'except:' catches all exceptions including SystemExit/KeyboardInterrupt - catch specific exception types instead",
            })
    return issues


# ─────────────────────────────────────────────────────────────────────────
#  Entry point
# ─────────────────────────────────────────────────────────────────────────

def _analyze_source(source):
    """Runs the full AST analysis pipeline on one file's source text.
    Returns a dict shaped like the single-file output contract:
    {'ok': True, 'issues': [...]} or {'ok': False, 'error': str}.
    Never raises - every failure mode is captured and reported in-band.
    """
    try:
        tree = ast.parse(source)
    except SyntaxError as e:
        return {'ok': False, 'error': 'SyntaxError: %s' % e}
    except Exception as e:  # pragma: no cover - defensive
        return {'ok': False, 'error': str(e)}

    issues = []

    try:
        importer = _UnusedImportChecker(tree)
        importer.collect()
        issues.extend(importer.unused())
    except Exception:
        pass

    try:
        for node in ast.walk(tree):
            if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
                issues.extend(_mutable_defaults(node))
                complexity = _compute_complexity(node)
                if complexity > COMPLEXITY_THRESHOLD:
                    issues.append({
                        'type': 'high_complexity',
                        'line': node.lineno,
                        'name': node.name,
                        'message': "Function '%s' has cyclomatic complexity %d (threshold %d)" % (
                            node.name, complexity, COMPLEXITY_THRESHOLD),
                    })
                issues.extend(_unused_locals(node))
    except Exception:
        pass

    try:
        dead_code_visitor = _DeadCodeVisitor()
        dead_code_visitor.visit(tree)
        issues.extend(dead_code_visitor.issues)
    except Exception:
        pass

    try:
        issues.extend(_find_bare_except(tree))
    except Exception:
        pass

    issues.sort(key=lambda i: i.get('line') or 0)
    return {'ok': True, 'issues': issues}


def _main_single():
    source = sys.stdin.read()
    print(json.dumps(_analyze_source(source)))


def _main_batch():
    """Batch mode: stdin is a JSON array of [{"path": str, "content": str}, ...].
    Analyzes every entry within this ONE interpreter invocation - avoids
    paying Python startup cost per file, which otherwise dominates scan
    time on Python-heavy codebases (see module docstring).
    """
    try:
        raw = sys.stdin.read()
        entries = json.loads(raw)
    except Exception as e:
        print(json.dumps({'ok': False, 'error': 'Failed to read batch input: %s' % e}))
        return

    results = []
    for entry in entries if isinstance(entries, list) else []:
        path = entry.get('path') if isinstance(entry, dict) else None
        content = entry.get('content') if isinstance(entry, dict) else None
        if content is None:
            results.append({'path': path, 'ok': False, 'issues': [], 'error': 'missing content'})
            continue
        outcome = _analyze_source(content)
        results.append({
            'path': path,
            'ok': outcome.get('ok', False),
            'issues': outcome.get('issues', []),
            'error': outcome.get('error'),
        })

    print(json.dumps({'ok': True, 'results': results}))


def main():
    if len(sys.argv) > 1 and sys.argv[1] == '--batch':
        _main_batch()
    else:
        _main_single()


if __name__ == '__main__':
    main()
