#!/usr/bin/env python3
# SPDX-License-Identifier: BSD-3-Clause
# Copyright (c) 2026 Robin Jarry

"""
Detect function local variable declarations that are not in reverse xmas
tree order (declaration lines sorted from longest to shortest).

Consecutive local variable declarations at the top of a block should be ordered
by descending source line length, e.g.:

    struct rte_ether_addr mac;
    struct nexthop *nh;
    uint16_t vrf_id;
    int ret;

The source is parsed with tree-sitter (no preprocessing, no headers, no build
required). Single-line declarations are grouped into maximal runs of
consecutive declaration statements, in every block including nested ones; an
empty line or a multi-line declaration (struct/union body, aggregate
initializer) starts a new run. Any declaration line longer than the one above
it in the same run is reported, unless its initializer references a variable
declared earlier in the run (it cannot be moved up).

Requires the tree-sitter python bindings and C grammar:

    pip install --user tree_sitter tree_sitter_c
"""

import argparse
import os
import re
import sys

try:
    import tree_sitter_c
    from tree_sitter import Language, Node, Parser
except ImportError:
    sys.exit("error: missing tree-sitter, run: pip install --user tree_sitter tree_sitter_c")

# Object-like macros that expand to nothing and only annotate declarations.
# tree-sitter does not preprocess, so they are blanked out before parsing to
# avoid confusing the declaration grammar (line lengths are measured on the
# original source).
EMPTY_MACROS = re.compile(rb"\b(?:vec)\b")

DECLARATORS = frozenset(
    (
        "identifier",
        "pointer_declarator",
        "array_declarator",
        "init_declarator",
        "function_declarator",
        "parenthesized_declarator",
    )
)

# grout indents with tabs; measure line lengths as they are displayed
TAB_WIDTH = 8


def visual_width(text: bytes) -> int:
    """Displayed width of a byte string, expanding tabs to TAB_WIDTH columns."""
    return len(text.decode(errors="replace").expandtabs(TAB_WIDTH))


def declared_name(node: Node) -> str:
    """Name of the variable a declarator introduces."""
    if node.type == "identifier":
        return node.text.decode()
    child = node.child_by_field_name("declarator")
    if child is not None:
        return declared_name(child)
    for child in node.named_children:
        if child.type in DECLARATORS or child.type == "identifier":
            return declared_name(child)
    return None


def declared_names(decl: Node) -> set[str]:
    """Names of all variables a declaration statement introduces."""
    names = set()
    for child in decl.named_children:
        if child.type in DECLARATORS:
            name = declared_name(child)
            if name is not None:
                names.add(name)
    return names


def initializer_names(decl: Node) -> set[str]:
    """Names of the variables referenced in a declaration's initializers."""
    names = set()

    def walk(node):
        if node.type == "identifier":
            names.add(node.text.decode())
        for child in node.named_children:
            walk(child)

    for child in decl.named_children:
        if child.type == "init_declarator":
            value = child.child_by_field_name("value")
            if value is not None:
                walk(value)
    return names


def measure(node: Node, lines: list[str]) -> int:
    """
    Displayed length of a single-line declaration, up to its ';' and excluding
    any trailing comment. Tabs count as TAB_WIDTH columns.
    """
    row, _ = node.start_point
    ecol = node.end_point[1]
    return visual_width(lines[row][:ecol])


def is_function_definition(node: Node) -> bool:
    """
    True for a real function definition (its declarator is a function
    declarator). tree-sitter misparses e.g. 'struct __rte_cache_aligned foo {
    ... }' as a function_definition whose declarator is a bare identifier; such
    a struct body must not be checked.
    """
    declarator = node.child_by_field_name("declarator")
    while declarator is not None and declarator.type in (
        "pointer_declarator",
        "parenthesized_declarator",
    ):
        declarator = declarator.child_by_field_name("declarator")
    return declarator is not None and declarator.type == "function_declarator"


def in_function_body(node: Node) -> bool:
    """
    True if a compound statement is a function-body block, not a brace block
    passed as a macro/function argument (tree-sitter parses the '{ ... }'
    argument of e.g. GR_IFACE_INFO(...) as a compound statement too).
    """
    parent = node.parent
    while parent is not None:
        if parent.type in ("argument_list", "call_expression"):
            return False
        if parent.type == "function_definition":
            return is_function_definition(parent)
        parent = parent.parent
    return False


def check_block(compound: Node, lines: list[str], report: callable):
    """
    Check every maximal run of consecutive variable declarations that are the
    direct children of a compound statement.
    """
    run = []

    def flush():
        earlier = set()
        for i, cur in enumerate(run):
            if i > 0:
                prev = run[i - 1]
                if measure(cur, lines) > measure(prev, lines) and (
                    initializer_names(cur).isdisjoint(earlier)
                ):
                    report(prev, cur)
            earlier |= declared_names(cur)
        run.clear()

    for child in compound.named_children:
        if child.type == "comment":
            continue
        if child.type == "declaration" and declared_names(child):
            if child.start_point[0] != child.end_point[0]:
                # multi-line declaration (struct/union body, aggregate
                # initializer): not part of the reverse xmas tree
                flush()
                continue
            if run:
                prev = run[-1]
                blank = any(
                    not lines[row].strip()
                    for row in range(prev.end_point[0] + 1, child.start_point[0])
                )
                if blank or child.start_point[0] <= prev.start_point[0]:
                    # empty line, macro or same-line trickery: start a new run
                    flush()
            run.append(child)
        else:
            flush()
    flush()


def check_file(parser: Parser, path: str, lines_cache: dict, report: callable):
    with open(path, "rb") as f:
        raw = f.read()
    lines = raw.splitlines()
    lines_cache[path] = lines
    source = EMPTY_MACROS.sub(lambda m: b" " * len(m.group()), raw)
    tree = parser.parse(source)
    stack = [tree.root_node]
    while stack:
        node = stack.pop()
        if node.type == "compound_statement" and in_function_body(node):
            check_block(node, lines, lambda p, c: report(path, p, c))
        stack.extend(node.named_children)


def main():
    parser = argparse.ArgumentParser(
        description=__doc__,
        formatter_class=argparse.RawDescriptionHelpFormatter,
    )
    parser.add_argument(
        "-C",
        "--force-color",
        action="store_true",
        help="Force colored output even if standard output is not a TTY.",
    )
    parser.add_argument(
        "files",
        nargs="+",
        help="Source files to check.",
    )
    args = parser.parse_args()

    color = args.force_color or sys.stdout.isatty() and "NO_COLOR" not in os.environ

    def paint(code, text):
        return f"\033[{code}m{text}\033[0m" if color else text

    lines_cache = {}
    violations = 0

    def report(path: str, prev: Node, cur: Node):
        nonlocal violations
        violations += 1
        lines = lines_cache[path]
        line = cur.start_point[0] + 1
        prev_len = measure(prev, lines)
        cur_len = measure(cur, lines)
        pstr = lines[prev.start_point[0]].decode(errors="replace").strip()
        craw = lines[cur.start_point[0]].decode(errors="replace")
        cstr = craw.strip()
        # the current declaration overflows the previous (shorter) one at the
        # end of its line, up to the ';' (comments are ignored)
        ncaret = cur_len - prev_len
        lead = craw[: len(craw) - len(craw.lstrip())]
        end = cur_len - len(lead.expandtabs(TAB_WIDTH))
        off = max(0, end - ncaret)
        print(
            f"{paint('35', path)}:{paint('36', line)}: "
            + paint(
                "31",
                "variable declaration not in reverse xmas tree order",
            )
        )
        print(f"    {pstr}")
        print(f"    {cstr[:off]}{paint('1;31', cstr[off:end])}{cstr[end:]}")
        print(f"    {' ' * off}{paint('1;33', '^' * ncaret)}")

    ts_parser = Parser(Language(tree_sitter_c.language()))
    for path in args.files:
        check_file(ts_parser, path, lines_cache, report)

    if violations:
        print(
            "\n"
            + paint(
                "1;31",
                f"{violations} declaration(s) not in reverse xmas tree order",
            )
        )
        return 1
    return 0


if __name__ == "__main__":
    try:
        sys.exit(main())
    except KeyboardInterrupt:
        print("interrupted", file=sys.stderr)
        sys.exit(130)
