slang-format.py

cross platform rendering playground

slang-format.py

5.79 KB
# LLM generated file
# based on https://github.com/shader-slang/slang/blob/master/source/slang/slang-language-server-auto-format.cpp
# slangd just silently does some pre/post processing for clang-format,
# this is just a tiny file to do the same. I just want some dedicated tool to format all shaders if needed

# single file
# python format_slang.py myshader.slang

# format directory recursively
# python format_slang.py ./src

# custom style
# python format_slang.py --style LLVM ./src

# multiple files
# python format_slang.py file1.slang file2.slang ./src


#!/usr/bin/env python3
import subprocess
import re
import os
from pathlib import Path
from xml.etree import ElementTree as ET


def find_clang_format():
    """Find clang-format executable"""
    # Try PATH first
    result = subprocess.run(
        ["where" if os.name == "nt" else "which", "clang-format"],
        capture_output=True,
        text=True,
    )
    if result.returncode == 0:
        return result.stdout.strip()
    return "clang-format"


def extract_spirv_asm_ranges(text):
    """Extract ranges of spirv_asm blocks to exclude from formatting"""
    ranges = []
    # Simple regex to find spirv_asm { ... } blocks
    pattern = r"spirv_asm\s*\{([^}]*(?:\{[^}]*\}[^}]*)*)\}"
    for match in re.finditer(pattern, text, re.DOTALL):
        ranges.append((match.start(), match.end()))
    return ranges


def format_source(text, clang_format_path, style, fallback_style):
    """Call clang-format and parse XML output"""

    # Build clang-format command
    cmd = [
        clang_format_path,
        "--assume-filename",
        "file.cs",
        "--output-replacements-xml",
    ]

    # Add style arguments
    if style.startswith("file"):
        # Check if .clang-format exists
        if os.path.exists(".clang-format") or os.path.exists("_clang_format"):
            cmd.extend(["-style", style])
        elif fallback_style:
            cmd.extend(["-style", fallback_style])
    elif style:
        cmd.extend(["-style", style])

    # Run clang-format with UTF-8 encoding
    result = subprocess.run(
        cmd,
        input=text,
        capture_output=True,
        text=True,
        encoding="utf-8",  # Add this
        errors="replace",  # Add this to handle any encoding issues gracefully
    )

    if result.returncode != 0:
        print(f"clang-format error: {result.stderr}")
        return []

    # Parse XML output
    try:
        root = ET.fromstring(result.stdout)
        edits = []
        for replacement in root.findall("replacement"):
            offset = int(replacement.get("offset", "0"))
            length = int(replacement.get("length", "0"))
            new_text = replacement.text or ""
            # Decode XML entities
            new_text = new_text.replace("
", "\r").replace("
", "\n")
            new_text = new_text.replace("&lt;", "<").replace("&gt;", ">")
            new_text = new_text.replace("&amp;", "&").replace("&apos;", "'")
            new_text = new_text.replace("&quot;", '"')
            edits.append((offset, length, new_text))
        return edits
    except ET.ParseError as e:
        print(f"Failed to parse clang-format output: {e}")
        return []


def apply_edits(text, edits, exclusion_ranges):
    """Apply text edits, respecting exclusion ranges"""
    # Filter out edits in spirv_asm blocks
    filtered_edits = []
    for offset, length, new_text in edits:
        skip = False
        for start, end in exclusion_ranges:
            if start <= offset < end:
                skip = True
                break
        if not skip:
            # Skip semicolon-after-brace special case
            if length == 0 and offset < len(text) and text[offset] == ";":
                if offset > 0 and text[offset - 1] == "}":
                    skip = True
            if not skip:
                filtered_edits.append((offset, length, new_text))

    # Apply edits in reverse order (to preserve offsets)
    result = text
    for offset, length, new_text in reversed(filtered_edits):
        result = result[:offset] + new_text + result[offset + length :]

    return result


def format_file(file_path, clang_format_path, style="file", fallback_style=None):
    """Format a single file"""
    with open(file_path, "r", encoding="utf-8") as f:
        content = f.read()

    exclusion_ranges = extract_spirv_asm_ranges(content)
    edits = format_source(content, clang_format_path, style, fallback_style)
    formatted = apply_edits(content, edits, exclusion_ranges)

    with open(file_path, "w", encoding="utf-8", newline="") as f:
        f.write(formatted)

    print(f"Formatted: {file_path}")


def main():
    import argparse

    parser = argparse.ArgumentParser(description="Format Slang files like the LSP does")
    parser.add_argument(
        "--style", default="file", help="clang-format style (default: file)"
    )
    parser.add_argument(
        "--fallback-style",
        default="{BasedOnStyle: Microsoft, BreakBeforeBraces: Allman, ColumnLimit: 0}",
        help="clang-format fallback style",
    )
    parser.add_argument(
        "--clang-format", default=None, help="Path to clang-format executable"
    )
    parser.add_argument("files", nargs="+", help="Files or directories to format")

    args = parser.parse_args()

    clang_format = args.clang_format or find_clang_format()
    print(f"Using clang-format: {clang_format}")

    # Gather files
    file_list = []
    for item in args.files:
        if os.path.isfile(item):
            file_list.append(item)
        elif os.path.isdir(item):
            file_list.extend(Path(item).rglob("*.slang"))

    # Format files
    for file_path in file_list:
        try:
            format_file(file_path, clang_format, args.style, args.fallback_style)
        except Exception as e:
            print(f"Error formatting {file_path}: {e}")


if __name__ == "__main__":
    main()