#!/usr/bin/env python3
"""Read Python files into code map facts, using Python's own parser.

This is a parser, not a model. Nothing here guesses: every object and every
link comes from the syntax tree, so the same file always produces the same
facts and two machines reading one project agree.

It runs on the person's own machine and prints facts, never source. What
leaves is "a function called price_of exists at line 40 of this file, and
these things call it". What it says is never read by anything but the parser.

Usage:  read_python.py <project-root> <relative-path-list-file>
Writes one JSON object to stdout.
"""
import ast
import json
import os
import sys

ROUTE_METHODS = {"get", "post", "put", "patch", "delete", "head", "options"}
MODEL_BASE_HINTS = {"Base", "BaseDataMixin", "DeclarativeBase", "Model",
                    "SQLModel", "models.Model"}
COLUMN_FACTORIES = {"Column", "mapped_column", "relationship"}


def dotted(node):
    """Render a Name / Attribute chain back to its source text."""
    if isinstance(node, ast.Name):
        return node.id
    if isinstance(node, ast.Attribute):
        base = dotted(node.value)
        return "%s.%s" % (base, node.attr) if base else node.attr
    if isinstance(node, ast.Call):
        return dotted(node.func)
    if isinstance(node, ast.Subscript):
        return dotted(node.value)
    return ""


def literal_str(node):
    if isinstance(node, ast.Constant) and isinstance(node.value, str):
        return node.value
    return None


class ModuleReader(ast.NodeVisitor):
    def __init__(self, file_id, module_path):
        self.file_id = file_id
        self.module_path = module_path
        self.objects = []
        self.links = []
        self.scope = []
        self.owner = [file_id]
        self.router_prefixes = {}
        self.mounts = []

    def qual(self, name):
        return ".".join(self.scope + [name])

    def add_object(self, object_id, kind, name, line, **extra):
        entry = {"id": object_id, "kind": kind, "name": name,
                 "file": self.file_id, "line": line}
        entry.update({k: v for k, v in extra.items() if v not in (None, [], "")})
        self.objects.append(entry)

    def add_link(self, src, kind, dst=None, to_name=None, **extra):
        link = {"from": src, "kind": kind}
        if dst:
            link["to"] = dst
        if to_name:
            link["to_name"] = to_name
        link.update(extra)
        self.links.append(link)

    # ---- imports ---------------------------------------------------------
    def visit_Import(self, node):
        for alias in node.names:
            self.add_link(self.file_id, "imports", to_name=alias.name, line=node.lineno)
        self.generic_visit(node)

    def visit_ImportFrom(self, node):
        module = node.module or ""
        if node.level:
            base = self.module_path.rsplit(".", node.level)[0]
            module = "%s.%s" % (base, module) if module else base
        for alias in node.names:
            self.add_link(
                self.file_id, "imports",
                to_name=("%s.%s" % (module, alias.name)) if module else alias.name,
                line=node.lineno)
        self.generic_visit(node)

    # ---- routers ---------------------------------------------------------
    def visit_Assign(self, node):
        """`router = APIRouter(prefix="/billing")` fixes every path below it."""
        if not self.scope and isinstance(node.value, ast.Call):
            leaf = dotted(node.value.func).split(".")[-1]
            if leaf in ("APIRouter", "Blueprint", "Router"):
                prefix = ""
                for kw in node.value.keywords:
                    if kw.arg in ("prefix", "url_prefix"):
                        prefix = literal_str(kw.value) or ""
                for target in node.targets:
                    if isinstance(target, ast.Name):
                        self.router_prefixes[target.id] = prefix
        self.generic_visit(node)

    def _include_router(self, node):
        """`include_router(billing.router, prefix="/api")` mounts a whole file."""
        prefix = ""
        for kw in node.keywords:
            if kw.arg in ("prefix", "url_prefix"):
                prefix = literal_str(kw.value) or ""
        if not node.args:
            return
        target = dotted(node.args[0])
        if target:
            self.mounts.append({"target": target, "prefix": prefix, "line": node.lineno})

    # ---- classes ---------------------------------------------------------
    def visit_ClassDef(self, node):
        qual = self.qual(node.name)
        object_id = "%s::%s" % (self.file_id, qual)
        bases = [dotted(b) for b in node.bases]
        tablename = None
        columns = []
        for stmt in node.body:
            if isinstance(stmt, ast.Assign):
                for target in stmt.targets:
                    if isinstance(target, ast.Name):
                        if target.id == "__tablename__":
                            tablename = literal_str(stmt.value)
                        elif isinstance(stmt.value, ast.Call) and \
                                dotted(stmt.value.func).split(".")[-1] in COLUMN_FACTORIES:
                            columns.append(target.id)
            elif isinstance(stmt, ast.AnnAssign) and isinstance(stmt.target, ast.Name):
                if isinstance(stmt.value, ast.Call) and \
                        dotted(stmt.value.func).split(".")[-1] in COLUMN_FACTORIES:
                    columns.append(stmt.target.id)

        is_model = bool(tablename) or any(
            b.split(".")[-1] in MODEL_BASE_HINTS for b in bases)
        self.add_object(object_id, "model" if is_model else "class", qual,
                        node.lineno, bases=bases, table=tablename, columns=columns)
        self.add_link(self.file_id, "defines", object_id)
        for base in bases:
            self.add_link(object_id, "inherits", to_name=base)
        if tablename:
            table_id = "table:%s" % tablename
            self.add_object(table_id, "table", tablename, node.lineno)
            self.add_link(object_id, "maps_to", table_id)

        self.scope.append(node.name)
        self.owner.append(object_id)
        self.generic_visit(node)
        self.owner.pop()
        self.scope.pop()

    # ---- functions -------------------------------------------------------
    def _function(self, node, is_async):
        qual = self.qual(node.name)
        object_id = "%s::%s" % (self.file_id, qual)
        decorators = [dotted(d) for d in node.decorator_list]
        self.add_object(object_id, "method" if self.scope else "function", qual,
                        node.lineno, decorators=decorators,
                        is_async=is_async or None,
                        args=[a.arg for a in node.args.args])
        self.add_link(self.file_id, "defines", object_id)

        # An endpoint: @router.get("/things"). The path here is RELATIVE to its
        # own router; the service puts the two halves back together.
        for dec in node.decorator_list:
            if not isinstance(dec, ast.Call):
                continue
            target = dotted(dec.func)
            verb = target.split(".")[-1].lower()
            if verb not in ROUTE_METHODS:
                continue
            path = literal_str(dec.args[0]) if dec.args else None
            if path is None:
                continue
            router_var = target.rsplit(".", 1)[0] if "." in target else ""
            route_id = "route:%s %s@%s" % (verb.upper(), path, self.file_id)
            self.add_object(route_id, "route", "%s %s" % (verb.upper(), path),
                            node.lineno, method=verb.upper(), path=path,
                            router_var=router_var)
            self.add_link(route_id, "handled_by", object_id)
            self.add_link(self.file_id, "declares", route_id)

        self.scope.append(node.name)
        self.owner.append(object_id)
        self.generic_visit(node)
        self.owner.pop()
        self.scope.pop()

    def visit_FunctionDef(self, node):
        self._function(node, False)

    def visit_AsyncFunctionDef(self, node):
        self._function(node, True)

    def visit_Name(self, node):
        """A capitalized name used in an expression is nearly always a class or
        a model being referenced, which is how database usage shows up."""
        if isinstance(node.ctx, ast.Load) and node.id[:1].isupper():
            self.add_link(self.owner[-1], "references", to_name=node.id, line=node.lineno)
        self.generic_visit(node)

    def visit_Call(self, node):
        target = dotted(node.func)
        if target.split(".")[-1] in ("include_router", "register_blueprint"):
            self._include_router(node)
        if target:
            leaf = target.split(".")[-1]
            if leaf and not leaf.startswith("_"):
                self.add_link(self.owner[-1], "calls", to_name=leaf,
                              full_name=target, line=node.lineno)
        self.generic_visit(node)


def module_path_for(rel_path):
    trimmed = rel_path[:-3]
    if trimmed.endswith("/__init__"):
        trimmed = trimmed[: -len("/__init__")]
    return trimmed.replace("/", ".")


def main():
    if len(sys.argv) < 3:
        sys.stderr.write("usage: read_python.py <project-root> <path-list-file>\n")
        return 2
    root = os.path.abspath(sys.argv[1])
    with open(sys.argv[2], "r", encoding="utf8") as handle:
        rels = [line.strip() for line in handle if line.strip()]

    facts = {}
    failures = []
    for rel_path in rels:
        full = os.path.join(root, rel_path)
        try:
            with open(full, "r", encoding="utf-8") as handle:
                source = handle.read()
            tree = ast.parse(source, filename=rel_path)
        except (SyntaxError, UnicodeDecodeError, OSError, ValueError) as exc:
            failures.append({"file": rel_path, "reason": type(exc).__name__})
            continue

        reader = ModuleReader(rel_path, module_path_for(rel_path))
        reader.visit(tree)
        entry = {
            "objects": [{
                "id": rel_path, "kind": "file",
                "name": os.path.basename(rel_path), "file": rel_path, "line": 1,
                "lang": "python", "lines": source.count("\n") + 1,
                "module": module_path_for(rel_path),
            }] + reader.objects,
            "links": reader.links,
        }
        if reader.router_prefixes or reader.mounts:
            entry["routers"] = {"prefixes": reader.router_prefixes,
                                "mounts": reader.mounts}
        facts[rel_path] = entry

    json.dump({"facts": facts, "failures": failures}, sys.stdout,
              separators=(",", ":"))
    return 0


if __name__ == "__main__":
    sys.exit(main())
