From f0bd6004aff79f3438a2e9c3a0ae43bc068c2be1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Mon, 28 Sep 2026 18:08:39 +0200 Subject: [PATCH 1/6] refactor: drop generated-code formatting from next, move formatter to eve.formatting Remove automatic formatting (black / clang-format) of generated code from gt4py.next and keep it only for gt4py.cartesian (and tests). - New leaf module `gt4py.eve.formatting` with best-effort `format_python_source` (black) and `format_cpp_source` (clang-format); both return the input unchanged when the tool is missing or fails. - Remove the old formatting API from `gt4py.eve.codegen`. - gt4py.next: no formatting; `SourceCodeSpec` loses `formatter_key`/`formatter_options`; `otf.artifacts.format_source` removed. - gt4py.cartesian: C++ formatted once in `BaseGTBackend._make_extension_sources`; Python backends use `eve.formatting`. - tach: only gt4py.cartesian may depend on `gt4py.eve.formatting`. - Dependencies: black and clang-format moved to the `cartesian` extra and the `test` dependency group. --- .../ADRs/next/0012-GridTools_Cpp_OTF_Steps.md | 2 +- pyproject.toml | 7 +- src/gt4py/cartesian/backend/dace_backend.py | 17 +-- src/gt4py/cartesian/backend/debug_backend.py | 4 +- src/gt4py/cartesian/backend/gtc_common.py | 6 +- src/gt4py/cartesian/backend/gtcpp_backend.py | 14 +- .../cartesian/backend/module_generator.py | 18 +-- src/gt4py/cartesian/backend/numpy_backend.py | 4 +- .../cartesian/gtc/gtcpp/gtcpp_codegen.py | 6 +- src/gt4py/eve/__init__.py | 2 +- src/gt4py/eve/codegen.py | 128 +----------------- src/gt4py/eve/formatting.py | 63 +++++++++ src/gt4py/next/otf/artifacts.py | 38 +----- src/gt4py/next/otf/binding/nanobind.py | 7 +- .../codegens/gtfn/gtfn_module.py | 19 ++- .../program_processors/runners/roundtrip.py | 17 +-- tach.toml | 7 + tests/eve_tests/unit_tests/test_formatting.py | 57 ++++++++ .../integration_tests/cases_utils.py | 2 +- .../binding_tests/test_cpp_interface.py | 60 +++----- .../dace_tests/test_dace_bindings.py | 4 +- uv.lock | 16 ++- 22 files changed, 209 insertions(+), 289 deletions(-) create mode 100644 src/gt4py/eve/formatting.py create mode 100644 tests/eve_tests/unit_tests/test_formatting.py diff --git a/docs/development/ADRs/next/0012-GridTools_Cpp_OTF_Steps.md b/docs/development/ADRs/next/0012-GridTools_Cpp_OTF_Steps.md index 4c7a132d24..77b9dabdc0 100644 --- a/docs/development/ADRs/next/0012-GridTools_Cpp_OTF_Steps.md +++ b/docs/development/ADRs/next/0012-GridTools_Cpp_OTF_Steps.md @@ -10,7 +10,7 @@ tags: [backend, cpp, gridtools, bindings, otf] - **Updated**: 2026-09-15 > [!NOTE] -> The entry points named below have since been renamed or moved: `program_processors.formatters.gtfn.format_sourcecode` is now `format_cpp`; `program_processors.codegens.gtfn_modules.translate_program` is now the `GTFNTranslationStep` class in `program_processors.codegens.gtfn.gtfn_module`; `otf.binding.pybind.bind_source` is now `otf.binding.nanobind.create_bindings` (nanobind replaced pybind11); and `program_processors.runners.gtfn_cpu.run_gtfn` is now `program_processors.runners.gtfn.run_gtfn`. `otf.step_types` was merged into `otf.stages`, `otf.workflow.Step` is now the `otf.workflow.Workflow` protocol, and `processor_interface` was split up — `ProgramFormatter` lives in `program_processors.program_formatter`. The decision — building the GTFN backend out of composable OTF steps — is unchanged. +> The entry points named below have since been renamed or moved: `program_processors.formatters.gtfn.format_sourcecode` is now `format_cpp`; `program_processors.codegens.gtfn_modules.translate_program` is now the `GTFNTranslationStep` class in `program_processors.codegens.gtfn.gtfn_module`; `otf.binding.pybind.bind_source` is now `otf.binding.nanobind.create_bindings` (nanobind replaced pybind11); and `program_processors.runners.gtfn_cpu.run_gtfn` is now `program_processors.runners.gtfn.run_gtfn`. `otf.step_types` was merged into `otf.stages`, `otf.workflow.Step` is now the `otf.workflow.Workflow` protocol, and `processor_interface` was split up — `ProgramFormatter` lives in `program_processors.program_formatter`. The decision — building the GTFN backend out of composable OTF steps — is unchanged. `SourceCodeSpec` no longer carries formatting options, and generated code is no longer auto-formatted in `gt4py.next`. ## Context diff --git a/pyproject.toml b/pyproject.toml index e5717fe037..c7db856b51 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -43,6 +43,8 @@ profiling = [ ] scripts = ["pyyaml>=6.0.1", "typer>=0.16.0", "packaging"] test = [ + 'black>=25.11', + 'clang-format>=18.1', 'hypothesis>=6.0.0', 'nbmake>=1.4.6', 'nox>=2025.02.09', @@ -90,7 +92,6 @@ classifiers = [ dependencies = [ 'attrs>=21.3', 'array-api-compat>=1.13', - 'black>=25.11', 'boltons>=20.1', 'cached-property>=1.5.1', 'click>=8.0.0', @@ -142,7 +143,7 @@ readme = 'README.md' requires-python = '>=3.12, <3.15, !=3.13.10, !=3.14.1' [project.optional-dependencies] -cartesian = ['gt4py[jax,standard,testing]'] +cartesian = ['black>=25.11', 'clang-format>=18.1', 'gt4py[jax,standard,testing]'] cuda12 = ['cupy-cuda12x>=12.0'] cuda13 = ['cupy-cuda13x>=14.0'] jax = [ @@ -154,7 +155,7 @@ jax-cuda13 = ['jax[cuda13-local]>=0.7.0', 'gt4py[cuda13]'] next = ['gt4py[jax,standard,testing]'] rocm6 = ['cupy>=13.4.1,<14.0'] rocm7 = ['cupy-rocm-7-0>=14.0'] -standard = ['clang-format>=18.1', 'scipy>=1.16.1'] +standard = ['scipy>=1.16.1'] testing = ['hypothesis>=6.93', 'pytest>=7.0'] [project.urls] diff --git a/src/gt4py/cartesian/backend/dace_backend.py b/src/gt4py/cartesian/backend/dace_backend.py index bd4b456349..38867a6946 100644 --- a/src/gt4py/cartesian/backend/dace_backend.py +++ b/src/gt4py/cartesian/backend/dace_backend.py @@ -44,7 +44,6 @@ from gt4py.cartesian.gtc.passes.oir_optimizations.utils import compute_fields_extents from gt4py.cartesian.gtc.passes.oir_pipeline import DefaultPipeline from gt4py.cartesian.utils import shash -from gt4py.eve import codegen from gt4py.eve.codegen import MakoTemplate as as_mako from gt4py.storage.cartesian import layout, layout_registry @@ -510,7 +509,7 @@ def __call__(self) -> dict[str, dict[str, str]]: implementation = DaCeComputationCodegen.apply(self.backend.builder, sdfg) - bindings = DaCeBindingsCodegen.apply(sdfg, self.module_name, backend=self.backend) + bindings = DaCeBindingsCodegen(self.backend).generate_sdfg_bindings(sdfg, self.module_name) bindings_ext = "cu" if self.backend.storage_info["device"] == "gpu" else "cpp" return { @@ -664,7 +663,7 @@ def apply(cls, builder: StencilBuilder, sdfg: SDFG) -> str: state_suffix=config.Config.get("compiler.codegen_state_struct_suffix"), ) computations = cls._postprocess_dace_code(code_objects, is_gpu) - generated_code = f"""\ + return f"""\ #include #include #include @@ -676,11 +675,6 @@ def apply(cls, builder: StencilBuilder, sdfg: SDFG) -> str: {interface} """ - if builder.options.format_source: - generated_code = codegen.format_source("cpp", generated_code, style="LLVM") - - return generated_code - def generate_dace_args(self, stencil_ir: gtir.Stencil, sdfg: SDFG) -> list[str]: oir = GTIRToOIR().visit(stencil_ir) field_extents = compute_fields_extents(oir, add_k=True) @@ -857,13 +851,6 @@ def generate_sdfg_bindings(self, sdfg: SDFG, module_name: str) -> str: sid_params=self.generate_sid_params(sdfg), ) - @classmethod - def apply(cls, sdfg: SDFG, module_name: str, *, backend: BaseDaceBackend) -> str: - generated_code = cls(backend).generate_sdfg_bindings(sdfg, module_name) - if backend.builder.options.format_source: - generated_code = codegen.format_source("cpp", generated_code, style="LLVM") - return generated_code - class DaCePyExtModuleGenerator(PyExtModuleGenerator): def __init__(self, builder: StencilBuilder) -> None: diff --git a/src/gt4py/cartesian/backend/debug_backend.py b/src/gt4py/cartesian/backend/debug_backend.py index 08d139d3f5..5892cabfad 100644 --- a/src/gt4py/cartesian/backend/debug_backend.py +++ b/src/gt4py/cartesian/backend/debug_backend.py @@ -16,7 +16,7 @@ from gt4py.cartesian.gtc.debug.debug_codegen import DebugCodeGen from gt4py.cartesian.gtc.gtir_to_oir import GTIRToOIR from gt4py.cartesian.gtc.passes import oir_optimizations -from gt4py.eve import codegen +from gt4py.eve import formatting from gt4py.storage import layout from gt4py.storage.cartesian import layout_registry @@ -45,7 +45,7 @@ def _generate_computation(self) -> dict[str, str | dict]: source_code = DebugCodeGen().visit(oir) if self.builder.options.format_source: - source_code = codegen.format_source("python", source_code) + source_code = formatting.format_python_source(source_code) caching = self.builder.caching computation_name = f"{caching.module_prefix}computation{caching.module_postfix}.py" diff --git a/src/gt4py/cartesian/backend/gtc_common.py b/src/gt4py/cartesian/backend/gtc_common.py index c8b110db9e..fff50e431e 100644 --- a/src/gt4py/cartesian/backend/gtc_common.py +++ b/src/gt4py/cartesian/backend/gtc_common.py @@ -19,7 +19,7 @@ from gt4py.cartesian.backend.module_generator import BaseModuleGenerator, ModuleData from gt4py.cartesian.gtc import gtir, utils as gtc_utils from gt4py.cartesian.gtc.passes.oir_pipeline import OirPipeline -from gt4py.eve import codegen +from gt4py.eve import codegen, formatting if TYPE_CHECKING: @@ -273,6 +273,10 @@ def _make_extension_sources(self) -> dict[str, dict[str, str]]: ) gt_pyext_generator = self.PYEXT_GENERATOR_CLASS(class_name, module_name, self) gt_pyext_sources = gt_pyext_generator() + if self.builder.options.format_source: + for sources in gt_pyext_sources.values(): + for file_name, source in sources.items(): + sources[file_name] = formatting.format_cpp_source(source) final_ext = ".cu" if self.languages and self.languages["computation"] == "cuda" else ".cpp" comp_src = gt_pyext_sources["computation"] for key in [k for k in comp_src.keys() if k.endswith(".src")]: diff --git a/src/gt4py/cartesian/backend/gtcpp_backend.py b/src/gt4py/cartesian/backend/gtcpp_backend.py index d77c2fb37f..25abd66e78 100644 --- a/src/gt4py/cartesian/backend/gtcpp_backend.py +++ b/src/gt4py/cartesian/backend/gtcpp_backend.py @@ -45,15 +45,11 @@ def __call__(self) -> dict[str, dict[str, str]]: ) oir_node = oir_pipeline.run(base_oir) gtcpp_ir = OIRToGTCpp().visit(oir_node) - format_source = self.backend.builder.options.format_source implementation = gtcpp_codegen.GTCppCodegen.apply( - gtcpp_ir, gt_backend_t=self.backend.GT_BACKEND_T, format_source=format_source + gtcpp_ir, gt_backend_t=self.backend.GT_BACKEND_T ) bindings = GTCppBindingsCodegen.apply( - gtcpp_ir, - module_name=self.module_name, - backend=self.backend, - format_source=format_source, + gtcpp_ir, module_name=self.module_name, backend=self.backend ) bindings_ext = ".cu" if self.backend.GT_BACKEND_T == "gpu" else ".cpp" return { @@ -115,11 +111,7 @@ def visit_Program(self, node: gtcpp.Program, **kwargs): @classmethod def apply(cls, root, *, module_name="stencil", **kwargs) -> str: - generated_code = cls(kwargs.get("backend")).visit(root, module_name=module_name, **kwargs) - if kwargs.get("format_source", True): - generated_code = codegen.format_source("cpp", generated_code, style="LLVM") - - return generated_code + return cls(kwargs.get("backend")).visit(root, module_name=module_name, **kwargs) class GTBaseBackend(BaseGTBackend): diff --git a/src/gt4py/cartesian/backend/module_generator.py b/src/gt4py/cartesian/backend/module_generator.py index 4ce5279feb..4ed0b1f48e 100644 --- a/src/gt4py/cartesian/backend/module_generator.py +++ b/src/gt4py/cartesian/backend/module_generator.py @@ -25,7 +25,7 @@ from gt4py.cartesian.gtc.passes.oir_access_kinds import compute_access_kinds from gt4py.cartesian.gtc.passes.oir_optimizations.utils import compute_fields_extents from gt4py.cartesian.gtc.utils import dimension_flags_to_names -from gt4py.eve import codegen +from gt4py.eve import formatting if TYPE_CHECKING: @@ -107,7 +107,6 @@ def make_args_data_from_gtir(pipeline: GtirPipeline) -> ModuleData: class BaseModuleGenerator(abc.ABC): - SOURCE_LINE_LENGTH = 120 TEMPLATE_INDENT_SIZE = 4 TEMPLATE_RESOURCE = "stencil_module.py.in" @@ -149,10 +148,8 @@ def __call__(self, args_data: ModuleData) -> str: post_run=self.generate_post_run(), implementation=self.generate_implementation(), ) - if self.builder.options.as_dict()["format_source"]: - module_source = codegen.format_source( - "python", module_source, line_length=self.SOURCE_LINE_LENGTH - ) + if self.builder.options.format_source: + module_source = formatting.format_python_source(module_source) return module_source @@ -202,16 +199,11 @@ def generate_backend_name(self) -> str: def generate_sources(self) -> dict[str, str]: """ - Return the source code of the stencil definition in string format. + Return the source code of the stencil definition verbatim, in string format. This is unlikely to require overriding. """ - if self.builder.gtir.sources is not None: - return { - key: codegen.format_source("python", value, line_length=self.SOURCE_LINE_LENGTH) - for key, value in self.builder.gtir.sources.items() - } - return {} + return dict(self.builder.gtir.sources or {}) def generate_constants(self) -> dict[str, str]: """ diff --git a/src/gt4py/cartesian/backend/numpy_backend.py b/src/gt4py/cartesian/backend/numpy_backend.py index c1c1c9e0b4..335feccf04 100644 --- a/src/gt4py/cartesian/backend/numpy_backend.py +++ b/src/gt4py/cartesian/backend/numpy_backend.py @@ -16,7 +16,7 @@ from gt4py.cartesian.gtc.gtir_to_oir import GTIRToOIR from gt4py.cartesian.gtc.numpy import npir from gt4py.cartesian.gtc.passes import oir_optimizations as oir_opt -from gt4py.eve import codegen +from gt4py.eve import formatting from gt4py.storage import layout from gt4py.storage.cartesian import layout_registry @@ -44,7 +44,7 @@ def generate_computation(self) -> dict[str, str | dict]: source = numpy.NpirCodegen.apply(self.npir, ignore_np_errstate=ignore_np_errstate) if self.builder.options.format_source: - source = codegen.format_source("python", source) + source = formatting.format_python_source(source) caching = self.builder.caching computation_name = f"{caching.module_prefix}computation{caching.module_postfix}.py" diff --git a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py index 62eae238a9..e22aab82c1 100644 --- a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py +++ b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py @@ -324,8 +324,4 @@ def apply(cls, root: LeafNode, **kwargs: Any) -> str: raise ValueError("apply() requires gtcpp.Progam root node") if "gt_backend_t" not in kwargs: raise TypeError("apply() missing 1 required keyword-only argument: 'gt_backend_t'") - generated_code = super().apply(root, offset_limit=_offset_limit(root), **kwargs) - if kwargs.get("format_source", True): - generated_code = codegen.format_source("cpp", generated_code, style="LLVM") - - return generated_code + return super().apply(root, offset_limit=_offset_limit(root), **kwargs) diff --git a/src/gt4py/eve/__init__.py b/src/gt4py/eve/__init__.py index 33295cefc6..70a4c999ad 100644 --- a/src/gt4py/eve/__init__.py +++ b/src/gt4py/eve/__init__.py @@ -11,7 +11,7 @@ The internal dependencies between modules are the following (each module depends on some of the previous ones): - 0. xtyping + 0. formatting, xtyping 1. exceptions, pattern_matching, type_definitions 2. utils 3. type_validation diff --git a/src/gt4py/eve/codegen.py b/src/gt4py/eve/codegen.py index e6e64adb55..e3ba40302e 100644 --- a/src/gt4py/eve/codegen.py +++ b/src/gt4py/eve/codegen.py @@ -14,17 +14,14 @@ import collections.abc import contextlib import inspect -import os import re import string -import subprocess import sys import textwrap import types -from collections.abc import Callable, Collection, Iterator, Mapping, Sequence +from collections.abc import Collection, Iterator, Mapping, Sequence from typing import Any, ClassVar, Optional, TypeVar, Union, overload -import black import jinja2 from mako import template as mako_tpl from typing_extensions import Protocol, runtime_checkable @@ -34,24 +31,6 @@ from .visitors import NodeVisitor -SourceFormatter = Callable[[str], str] - -SOURCE_FORMATTERS: dict[str, SourceFormatter] = {} -"""Global dict storing registered formatters.""" - - -class FormatterNameError(exceptions.EveRuntimeError): - """Run-time error registering a new source code formatter.""" - - ... - - -class FormattingError(exceptions.EveRuntimeError): - """Run-time error applying a source code formatter.""" - - ... - - class TemplateDefinitionError(exceptions.EveTypeError): """Template definition error.""" @@ -64,111 +43,6 @@ class TemplateRenderingError(exceptions.EveRuntimeError): ... -def register_formatter(language: str) -> Callable[[SourceFormatter], SourceFormatter]: - """Register source code formatters for specific languages (decorator).""" - - def _decorator(formatter: SourceFormatter) -> SourceFormatter: - if language in SOURCE_FORMATTERS: - raise FormatterNameError(f"Another formatter for language '{language}' already exists") - - assert callable(formatter) - SOURCE_FORMATTERS[language] = formatter - - return formatter - - return _decorator - - -@register_formatter("python") -def format_python_source( - source: str, - *, - line_length: int = 100, - python_versions: Optional[set[str]] = None, - string_normalization: bool = True, -) -> str: - """Format Python source code using black formatter.""" - python_versions = python_versions or {f"{sys.version_info.major}{sys.version_info.minor}"} - target_versions = set(black.TargetVersion[f"PY{v.replace('.', '')}"] for v in python_versions) # type: ignore[attr-defined] # .TargetVersion implicitly exported - - formatted_source = black.format_str( - source, - mode=black.FileMode( - line_length=line_length, - target_versions=target_versions, - string_normalization=string_normalization, - ), - ) - assert isinstance(formatted_source, str) - - return formatted_source - - -def _get_clang_format() -> Optional[str]: - """Return the clang-format executable, or None if not available.""" - executable = os.getenv("CLANG_FORMAT_EXECUTABLE", "clang-format") - try: - assert isinstance(executable, str) - if subprocess.run([executable, "--version"], capture_output=True).returncode != 0: - return None - except Exception: - return None - - return executable - - -_CLANG_FORMAT_EXECUTABLE = _get_clang_format() - - -if _CLANG_FORMAT_EXECUTABLE is not None: - - @register_formatter("cpp") - def format_cpp_source( - source: str, - *, - style: Optional[str] = None, - fallback_style: Optional[str] = None, - sort_includes: bool = False, - ) -> str: - """Format C++ source code using clang-format.""" - assert isinstance(_CLANG_FORMAT_EXECUTABLE, str) - args = [_CLANG_FORMAT_EXECUTABLE, "--assume-filename=_gt4py_generated_file.cpp"] - if style: - args.append(f"--style={style}") - if fallback_style: - args.append(f"--fallback-style={style}") - if sort_includes: - args.append("--sort-includes") - - try: - # use a timeout as clang-format used to deadlock on some sources - formatted_source = subprocess.run( - args, check=True, input=source, capture_output=True, text=True, timeout=3 - ).stdout - except subprocess.TimeoutExpired: - return source - - assert isinstance(formatted_source, str) - return formatted_source - - -def format_source(language: str, source: str, *, skip_errors: bool = True, **kwargs: Any) -> str: - """Format source code if a formatter exists for the specific language.""" - formatter = SOURCE_FORMATTERS.get(language, None) - try: - if formatter: - return formatter(source, **kwargs) - else: - raise FormattingError(f"Missing formatter for '{language}' language") - except Exception as e: - if skip_errors: - return source - else: - raise FormattingError( - f"Something went wrong when trying to format '{language}' source code" - ) from e - - class Name: """Text formatter with different case styles for symbol names in source code.""" diff --git a/src/gt4py/eve/formatting.py b/src/gt4py/eve/formatting.py new file mode 100644 index 0000000000..621627c1b5 --- /dev/null +++ b/src/gt4py/eve/formatting.py @@ -0,0 +1,63 @@ +# GT4Py - GridTools Framework +# +# Copyright (c) 2014-2024, ETH Zurich +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +"""Best-effort formatting of generated source code for human readers.""" + +from __future__ import annotations + +import os +import subprocess + + +def format_python_source(source: str) -> str: + """Format Python source code with `black`. + + Args: + source: Python source code. + + Returns: + The formatted source code, or `source` unchanged if `black` is not + installed or formatting fails. + """ + try: + # Lazy import: `black` is an optional dependency and slow to import. + import black + except ImportError: + return source + + try: + return black.format_str(source, mode=black.Mode(line_length=100)) + except ValueError: # `black.InvalidInput` for unparsable source + return source + + +def format_cpp_source(source: str) -> str: + """Format C++ source code with `clang-format` using the LLVM style. + + The executable can be overridden with the `CLANG_FORMAT_EXECUTABLE` + environment variable. + + Args: + source: C++ source code. + + Returns: + The formatted source code, or `source` unchanged if `clang-format` is + not available or formatting fails. + """ + args = [ + os.getenv("CLANG_FORMAT_EXECUTABLE", "clang-format"), + "--style=LLVM", + "--assume-filename=_gt4py_generated_file.cpp", + ] + try: + # use a timeout as clang-format used to deadlock on some sources + return subprocess.run( + args, input=source, capture_output=True, encoding="utf-8", check=True, timeout=3 + ).stdout + except (OSError, subprocess.SubprocessError, UnicodeError): + return source diff --git a/src/gt4py/next/otf/artifacts.py b/src/gt4py/next/otf/artifacts.py index 204e2e75d0..fa3226153b 100644 --- a/src/gt4py/next/otf/artifacts.py +++ b/src/gt4py/next/otf/artifacts.py @@ -25,28 +25,18 @@ from __future__ import annotations import dataclasses -import functools -from collections.abc import Callable, Mapping -from typing import Any, Generic, Optional, Protocol, TypeAlias, TypeVar, runtime_checkable +from collections.abc import Callable +from typing import Generic, Optional, Protocol, TypeAlias, TypeVar, runtime_checkable -from gt4py.eve import codegen from gt4py.next.otf.binding import interface @dataclasses.dataclass(frozen=True, kw_only=True) class SourceCodeSpec: - """ - Basic settings for any source programming language. - - Formatting will happen through ``eve.codegen.format_source``. - For available formatting options, check the options of the - specific formatter used depending on ``.formatter_key``. - """ + """Basic settings for any source programming language.""" source_language: str file_extension: str - formatter_key: str | None = None - formatter_options: Mapping[str, Any] | None = None @dataclasses.dataclass(frozen=True, kw_only=True) @@ -62,7 +52,6 @@ class PythonCodeSpec(SourceCodeSpec): source_language: str = "python" file_extension: str = "py" - formatter_key: str = "python" @dataclasses.dataclass(frozen=True, kw_only=True) @@ -85,10 +74,6 @@ class CPPCodeSpec(CPPLikeCodeSpec): source_language: str = "CXX" file_extension: str = "cpp" header_extension: str = "hpp" - formatter_key: str = "cpp" - formatter_options: Mapping[str, Any] = dataclasses.field( - default_factory=functools.partial(dict, style="LLVM") - ) @dataclasses.dataclass(frozen=True, kw_only=True) @@ -98,10 +83,6 @@ class CUDACodeSpec(CPPLikeCodeSpec): source_language: str = "CUDA" file_extension: str = "cu" header_extension: str = "cuh" - formatter_key: str = "cpp" - formatter_options: Mapping[str, Any] = dataclasses.field( - default_factory=functools.partial(dict, style="LLVM") - ) @dataclasses.dataclass(frozen=True, kw_only=True) @@ -111,19 +92,6 @@ class HIPCodeSpec(CPPLikeCodeSpec): source_language: str = "HIP" file_extension: str = "hip" header_extension: str = "h" - formatter_key: str = "cpp" - formatter_options: Mapping[str, Any] = dataclasses.field( - default_factory=functools.partial(dict, style="LLVM") - ) - - -def format_source(source_code_spec: SourceCodeSpec, source: str) -> str: - assert source_code_spec.formatter_key is not None, ( - "No formatter key specified in source code specification." - ) - return codegen.format_source( - source_code_spec.formatter_key, source, **(source_code_spec.formatter_options or {}) - ) CodeSpecT = TypeVar("CodeSpecT", bound=SourceCodeSpec) diff --git a/src/gt4py/next/otf/binding/nanobind.py b/src/gt4py/next/otf/binding/nanobind.py index ab4004d44e..f67b3383a6 100644 --- a/src/gt4py/next/otf/binding/nanobind.py +++ b/src/gt4py/next/otf/binding/nanobind.py @@ -309,12 +309,11 @@ def create_bindings( ), ) - src = artifacts.format_source( - program_source.code_spec, BindingCodeGenerator.apply(file_binding) + return artifacts.BindingSource( + BindingCodeGenerator.apply(file_binding), + (interface.LibraryDependency("nanobind", "2.0.0"),), ) - return artifacts.BindingSource(src, (interface.LibraryDependency("nanobind", "2.0.0"),)) - @dataclasses.dataclass(frozen=True) class ExtensionGenerator: diff --git a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py index dcf12cf509..ae8d75d03f 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -15,7 +15,6 @@ import numpy as np from gt4py._core import definitions as core_defs -from gt4py.eve import codegen from gt4py.next import common from gt4py.next.ffront import fbuiltins from gt4py.next.iterator import ir as itir @@ -171,8 +170,7 @@ def generate_stencil_source( column_axis=column_axis, ) - generated_code = GTFNCodegen.apply(gtfn_ir) - return codegen.format_source("cpp", generated_code, style="LLVM") + return GTFNCodegen.apply(gtfn_ir) def __call__( self, inp: stages.CompilableProgramDef @@ -209,14 +207,13 @@ def __call__( inp.args.offset_provider, inp.args.column_axis, ) - source_code = artifacts.format_source( - self._code_spec(), - f""" - #include <{self._backend_header()}> - #include - {stencil_src} - {decl_src} - """.strip(), + source_code = "\n".join( + [ + f"#include <{self._backend_header()}>", + "#include ", + stencil_src, + decl_src, + ] ) module: artifacts.ProgramSource[artifacts.HeaderAndSourceCodeSpec] = ( diff --git a/src/gt4py/next/program_processors/runners/roundtrip.py b/src/gt4py/next/program_processors/runners/roundtrip.py index 09f173d3f9..29afc3535a 100644 --- a/src/gt4py/next/program_processors/runners/roundtrip.py +++ b/src/gt4py/next/program_processors/runners/roundtrip.py @@ -110,8 +110,8 @@ def visit_Temporary(self, node: itir.Temporary, **kwargs: Any) -> str: # Caches the generated source by IR hash so re-codegen is skipped within a process. _SOURCE_CACHE: dict[int, tuple[str, str]] = {} -# Caches the loaded module by source string so re-exec is skipped within a process. -_MODULE_CACHE: dict[str, types.ModuleType] = {} +# Caches the loaded module by source string and debug flag so re-exec is skipped within a process. +_MODULE_CACHE: dict[tuple[str, bool], types.ModuleType] = {} def _generate_source( @@ -128,7 +128,6 @@ def _generate_source( ( ir, transforms, - debug, use_embedded, tuple(common.offset_provider_to_type(offset_provider).items()), ) @@ -142,11 +141,6 @@ def _generate_source( program = EmbeddedDSL.apply(ir) - # format output in debug mode for better debuggability - # (e.g. line numbers, overview in the debugger). - if debug: - program = codegen.format_python_source(program) - offset_literals: Iterable[str] = ( ir.pre_walk_values() .if_isinstance(itir.OffsetLiteral) @@ -187,8 +181,9 @@ def _generate_source( def _load_module(source_code: str, debug: bool) -> types.ModuleType: - if source_code in _MODULE_CACHE: - return _MODULE_CACHE[source_code] + cache_key = (source_code, debug) + if cache_key in _MODULE_CACHE: + return _MODULE_CACHE[cache_key] if debug: # Write to a real .py so debuggers/tracebacks have file/line info. @@ -205,7 +200,7 @@ def _load_module(source_code: str, debug: bool) -> types.ModuleType: mod = types.ModuleType("roundtrip_module") exec(compile(source_code, "", "exec"), mod.__dict__) - _MODULE_CACHE[source_code] = mod + _MODULE_CACHE[cache_key] = mod return mod diff --git a/tach.toml b/tach.toml index 78541c5dff..50abf97025 100644 --- a/tach.toml +++ b/tach.toml @@ -18,6 +18,7 @@ path = "gt4py.cartesian" depends_on = [ { path = "gt4py._core" }, { path = "gt4py.eve" }, + { path = "gt4py.eve.formatting" }, { path = "gt4py.storage" }, ] @@ -25,6 +26,12 @@ depends_on = [ path = "gt4py.eve" depends_on = [] +# Separate module so that only `gt4py.cartesian` can use the best-effort +# generated-code formatting (optional `black`/`clang-format` dependencies). +[[modules]] +path = "gt4py.eve.formatting" +depends_on = [] + [[modules]] path = "gt4py.next" depends_on = [ diff --git a/tests/eve_tests/unit_tests/test_formatting.py b/tests/eve_tests/unit_tests/test_formatting.py new file mode 100644 index 0000000000..2bb41147fd --- /dev/null +++ b/tests/eve_tests/unit_tests/test_formatting.py @@ -0,0 +1,57 @@ +# GT4Py - GridTools Framework +# +# Copyright (c) 2014-2024, ETH Zurich +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +import shutil +import sys + +import pytest + +from gt4py.eve import formatting + + +UNFORMATTED_PYTHON = "def f( a,b ):\n return a+b\n" +UNFORMATTED_CPP = "int f( int a,int b ){return a+b;}\n" + + +# -- Python tests -- +def test_format_python_source(): + pytest.importorskip("black") + assert formatting.format_python_source(UNFORMATTED_PYTHON) == ( + "def f(a, b):\n return a + b\n" + ) + + +def test_format_python_source_invalid_input(): + pytest.importorskip("black") + source = "def f(:\n" + assert formatting.format_python_source(source) == source + + +def test_format_python_source_without_black(monkeypatch): + monkeypatch.setitem(sys.modules, "black", None) + assert formatting.format_python_source(UNFORMATTED_PYTHON) == UNFORMATTED_PYTHON + + +# -- C++ tests -- +@pytest.mark.skipif(shutil.which("clang-format") is None, reason="clang-format not available") +def test_format_cpp_source(monkeypatch): + monkeypatch.delenv("CLANG_FORMAT_EXECUTABLE", raising=False) + assert ( + formatting.format_cpp_source(UNFORMATTED_CPP) == "int f(int a, int b) { return a + b; }\n" + ) + + +def test_format_cpp_source_missing_executable(monkeypatch): + monkeypatch.setenv("CLANG_FORMAT_EXECUTABLE", "/nonexistent") + assert formatting.format_cpp_source(UNFORMATTED_CPP) == UNFORMATTED_CPP + + +@pytest.mark.skipif(shutil.which("false") is None, reason="`false` executable not available") +def test_format_cpp_source_failing_executable(monkeypatch): + monkeypatch.setenv("CLANG_FORMAT_EXECUTABLE", "false") + assert formatting.format_cpp_source(UNFORMATTED_CPP) == UNFORMATTED_CPP diff --git a/tests/next_tests/integration_tests/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index 18617a6f91..3ea8821e9c 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -143,7 +143,7 @@ def debug_itir(tree): """Compare tree snippets while debugging.""" from devtools import debug - from gt4py.eve.codegen import format_python_source + from gt4py.eve.formatting import format_python_source from gt4py.next.program_processors import EmbeddedDSL debug(format_python_source(EmbeddedDSL.apply(tree))) diff --git a/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py b/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py index a25732649a..4daaaf074b 100644 --- a/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py +++ b/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py @@ -11,7 +11,7 @@ import gt4py.next as gtx import gt4py.next.otf.binding.cpp_interface as cpp import gt4py.next.type_system.type_specifications as ts -from gt4py.eve.codegen import format_source +from gt4py.eve.formatting import format_cpp_source from gt4py.next.otf.binding import interface @@ -27,28 +27,22 @@ def function_scalar_example(): def test_render_function_declaration_scalar(function_scalar_example): - rendered = format_source( - "cpp", cpp.render_function_declaration(function_scalar_example, "return;"), style="LLVM" + rendered = format_cpp_source( + cpp.render_function_declaration(function_scalar_example, "return;") ) - expected = format_source( - "cpp", - """\ + expected = format_cpp_source("""\ decltype(auto) example(double a, std::int64_t b) { return; }\ -""", - style="LLVM", - ) +""") assert rendered == expected def test_render_function_call_scalar(function_scalar_example): - rendered = format_source( - "cpp", - cpp.render_function_call(function_scalar_example, args=["13.6", "get_arg()"]), - style="LLVM", + rendered = format_cpp_source( + cpp.render_function_call(function_scalar_example, args=["13.6", "get_arg()"]) ) - expected = format_source("cpp", """example(13.6, get_arg())""", style="LLVM") + expected = format_cpp_source("""example(13.6, get_arg())""") assert rendered == expected @@ -75,29 +69,23 @@ def function_buffer_example(): def test_render_function_declaration_buffer(function_buffer_example): - rendered = format_source( - "cpp", cpp.render_function_declaration(function_buffer_example, "return;"), style="LLVM" + rendered = format_cpp_source( + cpp.render_function_declaration(function_buffer_example, "return;") ) - expected = format_source( - "cpp", - """\ + expected = format_cpp_source("""\ template decltype(auto) example(ArgT0&& a_buf, ArgT1&& b_buf) { return; }\ -""", - style="LLVM", - ) +""") assert rendered == expected def test_render_function_call_buffer(function_buffer_example): - rendered = format_source( - "cpp", - cpp.render_function_call(function_buffer_example, args=["get_arg_1()", "get_arg_2()"]), - style="LLVM", + rendered = format_cpp_source( + cpp.render_function_call(function_buffer_example, args=["get_arg_1()", "get_arg_2()"]) ) - expected = format_source("cpp", """example(get_arg_1(), get_arg_2())""", style="LLVM") + expected = format_cpp_source("""example(get_arg_1(), get_arg_2())""") assert rendered == expected @@ -126,25 +114,19 @@ def function_tuple_example(): def test_render_function_declaration_tuple(function_tuple_example): - rendered = format_source( - "cpp", cpp.render_function_declaration(function_tuple_example, "return;"), style="LLVM" - ) - expected = format_source( - "cpp", - """\ + rendered = format_cpp_source(cpp.render_function_declaration(function_tuple_example, "return;")) + expected = format_cpp_source("""\ template decltype(auto) example(ArgT0&& a_buf) { return; }\ -""", - style="LLVM", - ) +""") assert rendered == expected def test_render_function_call_tuple(function_tuple_example): - rendered = format_source( - "cpp", cpp.render_function_call(function_tuple_example, args=["get_arg_1()"]), style="LLVM" + rendered = format_cpp_source( + cpp.render_function_call(function_tuple_example, args=["get_arg_1()"]) ) - expected = format_source("cpp", """example(get_arg_1())""", style="LLVM") + expected = format_cpp_source("""example(get_arg_1())""") assert rendered == expected diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py index e4ec853c23..4b2d30ab43 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/dace_tests/test_dace_bindings.py @@ -12,7 +12,7 @@ import dace import numpy as np import pytest -from gt4py.eve import codegen +from gt4py.eve import formatting from gt4py import next as gtx from gt4py.next import common as gtx_common, int32 @@ -238,7 +238,7 @@ def mocked_compile_call( for line in inp.binding_source.source_code.splitlines() if not line.lstrip().startswith("assert") ) - assert codegen.format_python_source(binding_source_pruned) == binding_source_ref + assert formatting.format_python_source(binding_source_pruned) == binding_source_ref return _dace_compile_call(self, inp) diff --git a/uv.lock b/uv.lock index d024752ae3..e09ba514cd 100644 --- a/uv.lock +++ b/uv.lock @@ -1269,7 +1269,6 @@ source = { editable = "." } dependencies = [ { name = "array-api-compat" }, { name = "attrs" }, - { name = "black" }, { name = "boltons" }, { name = "cached-property" }, { name = "click" }, @@ -1301,6 +1300,7 @@ dependencies = [ [package.optional-dependencies] cartesian = [ + { name = "black" }, { name = "clang-format" }, { name = "hypothesis" }, { name = "jax" }, @@ -1325,7 +1325,6 @@ jax-cuda13 = [ { name = "jax", extra = ["cuda13-local"], marker = "extra == 'extra-5-gt4py-jax-cuda13' or (extra == 'extra-5-gt4py-cuda12' and extra == 'extra-5-gt4py-rocm6') or (extra == 'extra-5-gt4py-cuda12' and extra == 'extra-5-gt4py-rocm7') or (extra == 'extra-5-gt4py-rocm6' and extra == 'extra-5-gt4py-rocm7')" }, ] next = [ - { name = "clang-format" }, { name = "hypothesis" }, { name = "jax" }, { name = "pytest" }, @@ -1338,7 +1337,6 @@ rocm7 = [ { name = "cupy-rocm-7-0" }, ] standard = [ - { name = "clang-format" }, { name = "scipy" }, ] testing = [ @@ -1355,6 +1353,8 @@ build = [ ] dev = [ { name = "atlas4py" }, + { name = "black" }, + { name = "clang-format" }, { name = "cython" }, { name = "esbonio" }, { name = "hypothesis" }, @@ -1421,6 +1421,8 @@ scripts = [ { name = "typer" }, ] test = [ + { name = "black" }, + { name = "clang-format" }, { name = "hypothesis" }, { name = "nbmake" }, { name = "nox" }, @@ -1455,10 +1457,10 @@ typing-exports = [ requires-dist = [ { name = "array-api-compat", specifier = ">=1.13" }, { name = "attrs", specifier = ">=21.3" }, - { name = "black", specifier = ">=25.11" }, + { name = "black", marker = "extra == 'cartesian'", specifier = ">=25.11" }, { name = "boltons", specifier = ">=20.1" }, { name = "cached-property", specifier = ">=1.5.1" }, - { name = "clang-format", marker = "extra == 'standard'", specifier = ">=18.1" }, + { name = "clang-format", marker = "extra == 'cartesian'", specifier = ">=18.1" }, { name = "click", specifier = ">=8.0.0" }, { name = "cmake", specifier = ">=3.22" }, { name = "cupy", marker = "extra == 'rocm6'", specifier = ">=13.4.1,<14.0" }, @@ -1512,6 +1514,8 @@ build = [ ] dev = [ { name = "atlas4py", specifier = ">=0.41", index = "https://test.pypi.org/simple" }, + { name = "black", specifier = ">=25.11" }, + { name = "clang-format", specifier = ">=18.1" }, { name = "cython", specifier = ">=3.0.0" }, { name = "esbonio", specifier = ">=0.16.0" }, { name = "hypothesis", specifier = ">=6.0.0" }, @@ -1576,6 +1580,8 @@ scripts = [ { name = "typer", specifier = ">=0.16.0" }, ] test = [ + { name = "black", specifier = ">=25.11" }, + { name = "clang-format", specifier = ">=18.1" }, { name = "hypothesis", specifier = ">=6.0.0" }, { name = "nbmake", specifier = ">=1.4.6" }, { name = "nox", specifier = ">=2025.2.9" }, From b657b5dd27fd0d72c3b1963741a2bf57e51fa249 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Mon, 28 Sep 2026 19:01:53 +0200 Subject: [PATCH 2/6] test: cover generated-source formatting toggles and roundtrip debug module cache - cartesian: 'format_source' controls C++ formatting in 'BaseGTBackend._make_extension_sources'. - cartesian: stencil-definition sources are kept verbatim. - next: roundtrip module cache is keyed by (source, debug); only debug modules are backed by a real '.py' file. --- .../backend_tests/test_gtc_common.py | 36 +++++++++++++++++++ .../backend_tests/test_module_generator.py | 10 ++++++ .../runners_tests/test_roundtrip.py | 33 +++++++++++++++++ 3 files changed, 79 insertions(+) create mode 100644 tests/cartesian_tests/unit_tests/backend_tests/test_gtc_common.py create mode 100644 tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py diff --git a/tests/cartesian_tests/unit_tests/backend_tests/test_gtc_common.py b/tests/cartesian_tests/unit_tests/backend_tests/test_gtc_common.py new file mode 100644 index 0000000000..34ea013232 --- /dev/null +++ b/tests/cartesian_tests/unit_tests/backend_tests/test_gtc_common.py @@ -0,0 +1,36 @@ +# GT4Py - GridTools Framework +# +# Copyright (c) 2014-2024, ETH Zurich +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +import pytest + +from gt4py.cartesian.gtscript import PARALLEL, Field, computation, interval +from gt4py.cartesian.stencil_builder import StencilBuilder +from gt4py.eve import formatting + + +FORMATTED_MARK = "// formatted\n" + + +def sample_stencil(in_field: Field[float]): # type: ignore + with computation(PARALLEL), interval(...): # type: ignore + in_field += 1 # type: ignore + + +@pytest.mark.parametrize("format_source", [True, False]) +def test_make_extension_sources_formats_only_when_enabled(format_source, monkeypatch): + monkeypatch.setattr(formatting, "format_cpp_source", lambda source: FORMATTED_MARK + source) + builder = StencilBuilder(sample_stencil, backend="gt:cpu_ifirst").with_options( + name="sample_stencil", module=__name__, format_source=format_source + ) + + sources = builder.backend._make_extension_sources() + all_sources = [source for group in sources.values() for source in group.values()] + assert {"computation", "bindings"} <= sources.keys() + assert all_sources + assert all(source.startswith(FORMATTED_MARK) == format_source for source in all_sources) + assert not any(source.startswith(FORMATTED_MARK * 2) for source in all_sources) diff --git a/tests/cartesian_tests/unit_tests/backend_tests/test_module_generator.py b/tests/cartesian_tests/unit_tests/backend_tests/test_module_generator.py index 6478ecb10e..a097e36527 100644 --- a/tests/cartesian_tests/unit_tests/backend_tests/test_module_generator.py +++ b/tests/cartesian_tests/unit_tests/backend_tests/test_module_generator.py @@ -6,6 +6,8 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import types + import numpy as np import pytest @@ -62,6 +64,14 @@ def test_initialized_builder(sample_builder, sample_args_data): assert source +def test_generate_sources_is_verbatim(): + unformatted = "def f( x ):\n return x+1\n" + builder = types.SimpleNamespace(gtir=types.SimpleNamespace(sources={"f": unformatted})) + generator = SampleModuleGenerator(builder=builder) + + assert generator.generate_sources() == {"f": unformatted} + + def sample_stencil_with_args( used_io_field: Field[float], # type: ignore used_in_field: Field[float], # type: ignore diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py new file mode 100644 index 0000000000..a986a29048 --- /dev/null +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py @@ -0,0 +1,33 @@ +# GT4Py - GridTools Framework +# +# Copyright (c) 2014-2024, ETH Zurich +# All rights reserved. +# +# Please, refer to the LICENSE file in the root directory. +# SPDX-License-Identifier: BSD-3-Clause + +import os + +import pytest + +from gt4py.next.program_processors.runners import roundtrip + + +@pytest.mark.parametrize("debug_order", [(False, True), (True, False)]) +def test_load_module_caches_by_debug_flag(debug_order, monkeypatch): + monkeypatch.setattr(roundtrip, "_MODULE_CACHE", {}) + source_code = "VALUE = 42\n" + + modules = {debug: roundtrip._load_module(source_code, debug) for debug in debug_order} + try: + assert modules[False] is not modules[True] + assert modules[False].VALUE == modules[True].VALUE == 42 + # Only the debug module is backed by a real '.py' file. + assert not hasattr(modules[False], "__file__") + assert modules[True].__file__.endswith(".py") + assert os.path.isfile(modules[True].__file__) + + for debug in debug_order: + assert roundtrip._load_module(source_code, debug) is modules[debug] + finally: + os.remove(modules[True].__file__) From b32d9b44be9165ba3da770cdbf63affdef8c2d8e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Mon, 28 Sep 2026 21:30:52 +0200 Subject: [PATCH 3/6] refactor: keep gt4py.cartesian formatting behaviour unchanged gt4py.cartesian keeps its original formatting call sites, options and line lengths; it only switches from `eve.codegen.format_source` to `eve.formatting.format_cpp_source` / `format_python_source`. Drop the cartesian tests covering the reverted behaviour. `format_python_source` gains a `line_length` argument and pins black's target version to the running interpreter, as the old formatter did. --- src/gt4py/cartesian/backend/dace_backend.py | 17 +++++++-- src/gt4py/cartesian/backend/gtc_common.py | 6 +--- src/gt4py/cartesian/backend/gtcpp_backend.py | 16 ++++++--- .../cartesian/backend/module_generator.py | 16 ++++++--- .../cartesian/gtc/gtcpp/gtcpp_codegen.py | 8 +++-- src/gt4py/eve/formatting.py | 20 +++++++++-- .../backend_tests/test_gtc_common.py | 36 ------------------- .../backend_tests/test_module_generator.py | 10 ------ tests/eve_tests/unit_tests/test_formatting.py | 9 +++++ 9 files changed, 72 insertions(+), 66 deletions(-) delete mode 100644 tests/cartesian_tests/unit_tests/backend_tests/test_gtc_common.py diff --git a/src/gt4py/cartesian/backend/dace_backend.py b/src/gt4py/cartesian/backend/dace_backend.py index 38867a6946..f414f47a99 100644 --- a/src/gt4py/cartesian/backend/dace_backend.py +++ b/src/gt4py/cartesian/backend/dace_backend.py @@ -44,6 +44,7 @@ from gt4py.cartesian.gtc.passes.oir_optimizations.utils import compute_fields_extents from gt4py.cartesian.gtc.passes.oir_pipeline import DefaultPipeline from gt4py.cartesian.utils import shash +from gt4py.eve import formatting from gt4py.eve.codegen import MakoTemplate as as_mako from gt4py.storage.cartesian import layout, layout_registry @@ -509,7 +510,7 @@ def __call__(self) -> dict[str, dict[str, str]]: implementation = DaCeComputationCodegen.apply(self.backend.builder, sdfg) - bindings = DaCeBindingsCodegen(self.backend).generate_sdfg_bindings(sdfg, self.module_name) + bindings = DaCeBindingsCodegen.apply(sdfg, self.module_name, backend=self.backend) bindings_ext = "cu" if self.backend.storage_info["device"] == "gpu" else "cpp" return { @@ -663,7 +664,7 @@ def apply(cls, builder: StencilBuilder, sdfg: SDFG) -> str: state_suffix=config.Config.get("compiler.codegen_state_struct_suffix"), ) computations = cls._postprocess_dace_code(code_objects, is_gpu) - return f"""\ + generated_code = f"""\ #include #include #include @@ -675,6 +676,11 @@ def apply(cls, builder: StencilBuilder, sdfg: SDFG) -> str: {interface} """ + if builder.options.format_source: + generated_code = formatting.format_cpp_source(generated_code) + + return generated_code + def generate_dace_args(self, stencil_ir: gtir.Stencil, sdfg: SDFG) -> list[str]: oir = GTIRToOIR().visit(stencil_ir) field_extents = compute_fields_extents(oir, add_k=True) @@ -851,6 +857,13 @@ def generate_sdfg_bindings(self, sdfg: SDFG, module_name: str) -> str: sid_params=self.generate_sid_params(sdfg), ) + @classmethod + def apply(cls, sdfg: SDFG, module_name: str, *, backend: BaseDaceBackend) -> str: + generated_code = cls(backend).generate_sdfg_bindings(sdfg, module_name) + if backend.builder.options.format_source: + generated_code = formatting.format_cpp_source(generated_code) + return generated_code + class DaCePyExtModuleGenerator(PyExtModuleGenerator): def __init__(self, builder: StencilBuilder) -> None: diff --git a/src/gt4py/cartesian/backend/gtc_common.py b/src/gt4py/cartesian/backend/gtc_common.py index fff50e431e..c8b110db9e 100644 --- a/src/gt4py/cartesian/backend/gtc_common.py +++ b/src/gt4py/cartesian/backend/gtc_common.py @@ -19,7 +19,7 @@ from gt4py.cartesian.backend.module_generator import BaseModuleGenerator, ModuleData from gt4py.cartesian.gtc import gtir, utils as gtc_utils from gt4py.cartesian.gtc.passes.oir_pipeline import OirPipeline -from gt4py.eve import codegen, formatting +from gt4py.eve import codegen if TYPE_CHECKING: @@ -273,10 +273,6 @@ def _make_extension_sources(self) -> dict[str, dict[str, str]]: ) gt_pyext_generator = self.PYEXT_GENERATOR_CLASS(class_name, module_name, self) gt_pyext_sources = gt_pyext_generator() - if self.builder.options.format_source: - for sources in gt_pyext_sources.values(): - for file_name, source in sources.items(): - sources[file_name] = formatting.format_cpp_source(source) final_ext = ".cu" if self.languages and self.languages["computation"] == "cuda" else ".cpp" comp_src = gt_pyext_sources["computation"] for key in [k for k in comp_src.keys() if k.endswith(".src")]: diff --git a/src/gt4py/cartesian/backend/gtcpp_backend.py b/src/gt4py/cartesian/backend/gtcpp_backend.py index 25abd66e78..92305a3314 100644 --- a/src/gt4py/cartesian/backend/gtcpp_backend.py +++ b/src/gt4py/cartesian/backend/gtcpp_backend.py @@ -24,7 +24,7 @@ from gt4py.cartesian.gtc.gtcpp.oir_to_gtcpp import OIRToGTCpp from gt4py.cartesian.gtc.gtir_to_oir import GTIRToOIR from gt4py.cartesian.gtc.passes.oir_pipeline import DefaultPipeline -from gt4py.eve import codegen +from gt4py.eve import codegen, formatting from gt4py.storage.cartesian import layout, layout_registry @@ -45,11 +45,15 @@ def __call__(self) -> dict[str, dict[str, str]]: ) oir_node = oir_pipeline.run(base_oir) gtcpp_ir = OIRToGTCpp().visit(oir_node) + format_source = self.backend.builder.options.format_source implementation = gtcpp_codegen.GTCppCodegen.apply( - gtcpp_ir, gt_backend_t=self.backend.GT_BACKEND_T + gtcpp_ir, gt_backend_t=self.backend.GT_BACKEND_T, format_source=format_source ) bindings = GTCppBindingsCodegen.apply( - gtcpp_ir, module_name=self.module_name, backend=self.backend + gtcpp_ir, + module_name=self.module_name, + backend=self.backend, + format_source=format_source, ) bindings_ext = ".cu" if self.backend.GT_BACKEND_T == "gpu" else ".cpp" return { @@ -111,7 +115,11 @@ def visit_Program(self, node: gtcpp.Program, **kwargs): @classmethod def apply(cls, root, *, module_name="stencil", **kwargs) -> str: - return cls(kwargs.get("backend")).visit(root, module_name=module_name, **kwargs) + generated_code = cls(kwargs.get("backend")).visit(root, module_name=module_name, **kwargs) + if kwargs.get("format_source", True): + generated_code = formatting.format_cpp_source(generated_code) + + return generated_code class GTBaseBackend(BaseGTBackend): diff --git a/src/gt4py/cartesian/backend/module_generator.py b/src/gt4py/cartesian/backend/module_generator.py index 4ed0b1f48e..58f5c4f3c5 100644 --- a/src/gt4py/cartesian/backend/module_generator.py +++ b/src/gt4py/cartesian/backend/module_generator.py @@ -107,6 +107,7 @@ def make_args_data_from_gtir(pipeline: GtirPipeline) -> ModuleData: class BaseModuleGenerator(abc.ABC): + SOURCE_LINE_LENGTH = 120 TEMPLATE_INDENT_SIZE = 4 TEMPLATE_RESOURCE = "stencil_module.py.in" @@ -148,8 +149,10 @@ def __call__(self, args_data: ModuleData) -> str: post_run=self.generate_post_run(), implementation=self.generate_implementation(), ) - if self.builder.options.format_source: - module_source = formatting.format_python_source(module_source) + if self.builder.options.as_dict()["format_source"]: + module_source = formatting.format_python_source( + module_source, line_length=self.SOURCE_LINE_LENGTH + ) return module_source @@ -199,11 +202,16 @@ def generate_backend_name(self) -> str: def generate_sources(self) -> dict[str, str]: """ - Return the source code of the stencil definition verbatim, in string format. + Return the source code of the stencil definition in string format. This is unlikely to require overriding. """ - return dict(self.builder.gtir.sources or {}) + if self.builder.gtir.sources is not None: + return { + key: formatting.format_python_source(value, line_length=self.SOURCE_LINE_LENGTH) + for key, value in self.builder.gtir.sources.items() + } + return {} def generate_constants(self) -> dict[str, str]: """ diff --git a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py index e22aab82c1..5f0f30468d 100644 --- a/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py +++ b/src/gt4py/cartesian/gtc/gtcpp/gtcpp_codegen.py @@ -20,7 +20,7 @@ UnaryOperator, ) from gt4py.cartesian.gtc.gtcpp import gtcpp -from gt4py.eve import codegen +from gt4py.eve import codegen, formatting from gt4py.eve.codegen import FormatTemplate as as_fmt, MakoTemplate as as_mako from gt4py.eve.concepts import LeafNode @@ -324,4 +324,8 @@ def apply(cls, root: LeafNode, **kwargs: Any) -> str: raise ValueError("apply() requires gtcpp.Progam root node") if "gt_backend_t" not in kwargs: raise TypeError("apply() missing 1 required keyword-only argument: 'gt_backend_t'") - return super().apply(root, offset_limit=_offset_limit(root), **kwargs) + generated_code = super().apply(root, offset_limit=_offset_limit(root), **kwargs) + if kwargs.get("format_source", True): + generated_code = formatting.format_cpp_source(generated_code) + + return generated_code diff --git a/src/gt4py/eve/formatting.py b/src/gt4py/eve/formatting.py index 621627c1b5..ab4ae24c27 100644 --- a/src/gt4py/eve/formatting.py +++ b/src/gt4py/eve/formatting.py @@ -12,17 +12,22 @@ import os import subprocess +import sys -def format_python_source(source: str) -> str: +def format_python_source(source: str, *, line_length: int = 100) -> str: """Format Python source code with `black`. + The target Python version is pinned to the running interpreter. + Args: source: Python source code. + line_length: Maximum line length of the formatted code. Returns: The formatted source code, or `source` unchanged if `black` is not - installed or formatting fails. + installed, does not support the running interpreter, or formatting + fails. """ try: # Lazy import: `black` is an optional dependency and slow to import. @@ -31,7 +36,16 @@ def format_python_source(source: str) -> str: return source try: - return black.format_str(source, mode=black.Mode(line_length=100)) + target_version = black.TargetVersion[ # type: ignore[attr-defined] # .TargetVersion implicitly exported + f"PY{sys.version_info.major}{sys.version_info.minor}" + ] + except KeyError: # `black` is too old to know the running interpreter + return source + + try: + return black.format_str( + source, mode=black.Mode(line_length=line_length, target_versions={target_version}) + ) except ValueError: # `black.InvalidInput` for unparsable source return source diff --git a/tests/cartesian_tests/unit_tests/backend_tests/test_gtc_common.py b/tests/cartesian_tests/unit_tests/backend_tests/test_gtc_common.py deleted file mode 100644 index 34ea013232..0000000000 --- a/tests/cartesian_tests/unit_tests/backend_tests/test_gtc_common.py +++ /dev/null @@ -1,36 +0,0 @@ -# GT4Py - GridTools Framework -# -# Copyright (c) 2014-2024, ETH Zurich -# All rights reserved. -# -# Please, refer to the LICENSE file in the root directory. -# SPDX-License-Identifier: BSD-3-Clause - -import pytest - -from gt4py.cartesian.gtscript import PARALLEL, Field, computation, interval -from gt4py.cartesian.stencil_builder import StencilBuilder -from gt4py.eve import formatting - - -FORMATTED_MARK = "// formatted\n" - - -def sample_stencil(in_field: Field[float]): # type: ignore - with computation(PARALLEL), interval(...): # type: ignore - in_field += 1 # type: ignore - - -@pytest.mark.parametrize("format_source", [True, False]) -def test_make_extension_sources_formats_only_when_enabled(format_source, monkeypatch): - monkeypatch.setattr(formatting, "format_cpp_source", lambda source: FORMATTED_MARK + source) - builder = StencilBuilder(sample_stencil, backend="gt:cpu_ifirst").with_options( - name="sample_stencil", module=__name__, format_source=format_source - ) - - sources = builder.backend._make_extension_sources() - all_sources = [source for group in sources.values() for source in group.values()] - assert {"computation", "bindings"} <= sources.keys() - assert all_sources - assert all(source.startswith(FORMATTED_MARK) == format_source for source in all_sources) - assert not any(source.startswith(FORMATTED_MARK * 2) for source in all_sources) diff --git a/tests/cartesian_tests/unit_tests/backend_tests/test_module_generator.py b/tests/cartesian_tests/unit_tests/backend_tests/test_module_generator.py index a097e36527..6478ecb10e 100644 --- a/tests/cartesian_tests/unit_tests/backend_tests/test_module_generator.py +++ b/tests/cartesian_tests/unit_tests/backend_tests/test_module_generator.py @@ -6,8 +6,6 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause -import types - import numpy as np import pytest @@ -64,14 +62,6 @@ def test_initialized_builder(sample_builder, sample_args_data): assert source -def test_generate_sources_is_verbatim(): - unformatted = "def f( x ):\n return x+1\n" - builder = types.SimpleNamespace(gtir=types.SimpleNamespace(sources={"f": unformatted})) - generator = SampleModuleGenerator(builder=builder) - - assert generator.generate_sources() == {"f": unformatted} - - def sample_stencil_with_args( used_io_field: Field[float], # type: ignore used_in_field: Field[float], # type: ignore diff --git a/tests/eve_tests/unit_tests/test_formatting.py b/tests/eve_tests/unit_tests/test_formatting.py index 2bb41147fd..cd61699a2c 100644 --- a/tests/eve_tests/unit_tests/test_formatting.py +++ b/tests/eve_tests/unit_tests/test_formatting.py @@ -26,6 +26,15 @@ def test_format_python_source(): ) +def test_format_python_source_line_length(): + pytest.importorskip("black") + source = "result = function_name(first_argument, second_argument)\n" + assert formatting.format_python_source(source) == source + assert formatting.format_python_source(source, line_length=40) == ( + "result = function_name(\n first_argument, second_argument\n)\n" + ) + + def test_format_python_source_invalid_input(): pytest.importorskip("black") source = "def f(:\n" From af4457a1a211fa9f0037f54c6fa6373097bac1df Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Thu, 1 Oct 2026 15:30:01 +0200 Subject: [PATCH 4/6] fix[eve]: return Python source unchanged on any black failure Restores the catch-all of the removed codegen.format_source(skip_errors=True) path used by gt4py.cartesian. Adds tests for unexpected black errors, an interpreter unknown to black, and the debug-independent roundtrip source cache. --- src/gt4py/eve/formatting.py | 2 +- tests/eve_tests/unit_tests/test_formatting.py | 17 ++++++++ .../runners_tests/test_roundtrip.py | 39 +++++++++++++++++++ 3 files changed, 57 insertions(+), 1 deletion(-) diff --git a/src/gt4py/eve/formatting.py b/src/gt4py/eve/formatting.py index ab4ae24c27..a77258671a 100644 --- a/src/gt4py/eve/formatting.py +++ b/src/gt4py/eve/formatting.py @@ -46,7 +46,7 @@ def format_python_source(source: str, *, line_length: int = 100) -> str: return black.format_str( source, mode=black.Mode(line_length=line_length, target_versions={target_version}) ) - except ValueError: # `black.InvalidInput` for unparsable source + except Exception: # e.g. `black.InvalidInput` for unparsable source, or internal `black` errors return source diff --git a/tests/eve_tests/unit_tests/test_formatting.py b/tests/eve_tests/unit_tests/test_formatting.py index cd61699a2c..630bf834c2 100644 --- a/tests/eve_tests/unit_tests/test_formatting.py +++ b/tests/eve_tests/unit_tests/test_formatting.py @@ -41,6 +41,23 @@ def test_format_python_source_invalid_input(): assert formatting.format_python_source(source) == source +def test_format_python_source_black_failure(monkeypatch): + black = pytest.importorskip("black") + + def failing_format_str(*args, **kwargs): + raise RuntimeError("internal black error") + + monkeypatch.setattr(black, "format_str", failing_format_str) + assert formatting.format_python_source(UNFORMATTED_PYTHON) == UNFORMATTED_PYTHON + + +def test_format_python_source_unsupported_interpreter(monkeypatch): + black = pytest.importorskip("black") + # Simulate a `black` version that does not know the running interpreter. + monkeypatch.setattr(black, "TargetVersion", {}) + assert formatting.format_python_source(UNFORMATTED_PYTHON) == UNFORMATTED_PYTHON + + def test_format_python_source_without_black(monkeypatch): monkeypatch.setitem(sys.modules, "black", None) assert formatting.format_python_source(UNFORMATTED_PYTHON) == UNFORMATTED_PYTHON diff --git a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py index a986a29048..510aa1ef1d 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py @@ -10,9 +10,48 @@ import pytest +from gt4py.next.iterator import ir as itir +from gt4py.next.iterator.ir_utils import ir_makers as im from gt4py.next.program_processors.runners import roundtrip +@pytest.mark.parametrize("debug_order", [(False, True), (True, False)]) +def test_generate_source_ignores_debug_flag(debug_order, monkeypatch): + monkeypatch.setattr(roundtrip, "_SOURCE_CACHE", {}) + domain = im.call("cartesian_domain")(im.named_range(itir.AxisLiteral(value="D"), 0, 1)) + program = itir.Program( + id="testee", + function_definitions=[], + params=[itir.Sym(id="out")], + declarations=[], + body=[ + itir.SetAt(expr=im.as_fieldop("deref")(), domain=domain, target=itir.SymRef(id="out")) + ], + ) + transform_calls = [] + + def transforms(ir, *, offset_provider): + transform_calls.append(ir) + return ir + + sources = [ + roundtrip._generate_source( + program, + debug=debug, + use_embedded=True, + offset_provider={}, + transforms=transforms, + ) + for debug in debug_order + ] + + assert sources[0] == sources[1] + assert sources[0][1] == "testee" + # The source is generated once and shared by both debug modes. + assert len(transform_calls) == 1 + assert len(roundtrip._SOURCE_CACHE) == 1 + + @pytest.mark.parametrize("debug_order", [(False, True), (True, False)]) def test_load_module_caches_by_debug_flag(debug_order, monkeypatch): monkeypatch.setattr(roundtrip, "_MODULE_CACHE", {}) From 62519c455f8f68b70d2d728e0469a109c97e1e41 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 2 Oct 2026 10:28:55 +0200 Subject: [PATCH 5/6] refactor: address review comments - Import the gt4py.eve.formatting module instead of its functions in next test helpers (cases_utils.debug_itir, test_cpp_interface). --- .../integration_tests/cases_utils.py | 4 +-- .../binding_tests/test_cpp_interface.py | 28 ++++++++++--------- 2 files changed, 17 insertions(+), 15 deletions(-) diff --git a/tests/next_tests/integration_tests/cases_utils.py b/tests/next_tests/integration_tests/cases_utils.py index 3ea8821e9c..ea7c898f4d 100644 --- a/tests/next_tests/integration_tests/cases_utils.py +++ b/tests/next_tests/integration_tests/cases_utils.py @@ -143,10 +143,10 @@ def debug_itir(tree): """Compare tree snippets while debugging.""" from devtools import debug - from gt4py.eve.formatting import format_python_source + from gt4py.eve import formatting from gt4py.next.program_processors import EmbeddedDSL - debug(format_python_source(EmbeddedDSL.apply(tree))) + debug(formatting.format_python_source(EmbeddedDSL.apply(tree))) DimsType = TypeVar("DimsType") diff --git a/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py b/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py index 4daaaf074b..8e5e8dcca4 100644 --- a/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py +++ b/tests/next_tests/unit_tests/otf_tests/binding_tests/test_cpp_interface.py @@ -11,7 +11,7 @@ import gt4py.next as gtx import gt4py.next.otf.binding.cpp_interface as cpp import gt4py.next.type_system.type_specifications as ts -from gt4py.eve.formatting import format_cpp_source +from gt4py.eve import formatting from gt4py.next.otf.binding import interface @@ -27,10 +27,10 @@ def function_scalar_example(): def test_render_function_declaration_scalar(function_scalar_example): - rendered = format_cpp_source( + rendered = formatting.format_cpp_source( cpp.render_function_declaration(function_scalar_example, "return;") ) - expected = format_cpp_source("""\ + expected = formatting.format_cpp_source("""\ decltype(auto) example(double a, std::int64_t b) { return; }\ @@ -39,10 +39,10 @@ def test_render_function_declaration_scalar(function_scalar_example): def test_render_function_call_scalar(function_scalar_example): - rendered = format_cpp_source( + rendered = formatting.format_cpp_source( cpp.render_function_call(function_scalar_example, args=["13.6", "get_arg()"]) ) - expected = format_cpp_source("""example(13.6, get_arg())""") + expected = formatting.format_cpp_source("""example(13.6, get_arg())""") assert rendered == expected @@ -69,10 +69,10 @@ def function_buffer_example(): def test_render_function_declaration_buffer(function_buffer_example): - rendered = format_cpp_source( + rendered = formatting.format_cpp_source( cpp.render_function_declaration(function_buffer_example, "return;") ) - expected = format_cpp_source("""\ + expected = formatting.format_cpp_source("""\ template decltype(auto) example(ArgT0&& a_buf, ArgT1&& b_buf) { return; @@ -82,10 +82,10 @@ def test_render_function_declaration_buffer(function_buffer_example): def test_render_function_call_buffer(function_buffer_example): - rendered = format_cpp_source( + rendered = formatting.format_cpp_source( cpp.render_function_call(function_buffer_example, args=["get_arg_1()", "get_arg_2()"]) ) - expected = format_cpp_source("""example(get_arg_1(), get_arg_2())""") + expected = formatting.format_cpp_source("""example(get_arg_1(), get_arg_2())""") assert rendered == expected @@ -114,8 +114,10 @@ def function_tuple_example(): def test_render_function_declaration_tuple(function_tuple_example): - rendered = format_cpp_source(cpp.render_function_declaration(function_tuple_example, "return;")) - expected = format_cpp_source("""\ + rendered = formatting.format_cpp_source( + cpp.render_function_declaration(function_tuple_example, "return;") + ) + expected = formatting.format_cpp_source("""\ template decltype(auto) example(ArgT0&& a_buf) { return; @@ -125,8 +127,8 @@ def test_render_function_declaration_tuple(function_tuple_example): def test_render_function_call_tuple(function_tuple_example): - rendered = format_cpp_source( + rendered = formatting.format_cpp_source( cpp.render_function_call(function_tuple_example, args=["get_arg_1()"]) ) - expected = format_cpp_source("""example(get_arg_1())""") + expected = formatting.format_cpp_source("""example(get_arg_1())""") assert rendered == expected From 54b9caf8c1e1bffe1febfc709773e4538338bf02 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Enrique=20Gonz=C3=A1lez=20Paredes?= Date: Fri, 2 Oct 2026 10:55:25 +0200 Subject: [PATCH 6/6] refactor[eve]: probe clang-format availability once before formatting Mirror the previous eve.codegen behaviour: check CLANG_FORMAT_EXECUTABLE once with --version (lazily, cached) and skip C++ formatting when the tool is unavailable. Addresses review comment on PR #2927. --- src/gt4py/eve/formatting.py | 29 +++++++++++++++---- tests/eve_tests/unit_tests/test_formatting.py | 21 ++++++++++++++ 2 files changed, 44 insertions(+), 6 deletions(-) diff --git a/src/gt4py/eve/formatting.py b/src/gt4py/eve/formatting.py index a77258671a..4ba9612696 100644 --- a/src/gt4py/eve/formatting.py +++ b/src/gt4py/eve/formatting.py @@ -10,6 +10,7 @@ from __future__ import annotations +import functools import os import subprocess import sys @@ -54,7 +55,7 @@ def format_cpp_source(source: str) -> str: """Format C++ source code with `clang-format` using the LLVM style. The executable can be overridden with the `CLANG_FORMAT_EXECUTABLE` - environment variable. + environment variable. Its availability is checked once, on first use. Args: source: C++ source code. @@ -63,11 +64,11 @@ def format_cpp_source(source: str) -> str: The formatted source code, or `source` unchanged if `clang-format` is not available or formatting fails. """ - args = [ - os.getenv("CLANG_FORMAT_EXECUTABLE", "clang-format"), - "--style=LLVM", - "--assume-filename=_gt4py_generated_file.cpp", - ] + executable = _get_clang_format() + if executable is None: + return source + + args = [executable, "--style=LLVM", "--assume-filename=_gt4py_generated_file.cpp"] try: # use a timeout as clang-format used to deadlock on some sources return subprocess.run( @@ -75,3 +76,19 @@ def format_cpp_source(source: str) -> str: ).stdout except (OSError, subprocess.SubprocessError, UnicodeError): return source + + +@functools.cache +def _get_clang_format() -> str | None: + """Return the `clang-format` executable, or `None` if it is not available. + + The result is cached, so the executable is probed only once. + """ + executable = os.getenv("CLANG_FORMAT_EXECUTABLE", "clang-format") + try: + if subprocess.run([executable, "--version"], capture_output=True).returncode != 0: + return None + except Exception: + return None + + return executable diff --git a/tests/eve_tests/unit_tests/test_formatting.py b/tests/eve_tests/unit_tests/test_formatting.py index 630bf834c2..2fbe0970e7 100644 --- a/tests/eve_tests/unit_tests/test_formatting.py +++ b/tests/eve_tests/unit_tests/test_formatting.py @@ -64,6 +64,13 @@ def test_format_python_source_without_black(monkeypatch): # -- C++ tests -- +@pytest.fixture(autouse=True) +def clear_clang_format_cache(): + formatting._get_clang_format.cache_clear() + yield + formatting._get_clang_format.cache_clear() + + @pytest.mark.skipif(shutil.which("clang-format") is None, reason="clang-format not available") def test_format_cpp_source(monkeypatch): monkeypatch.delenv("CLANG_FORMAT_EXECUTABLE", raising=False) @@ -81,3 +88,17 @@ def test_format_cpp_source_missing_executable(monkeypatch): def test_format_cpp_source_failing_executable(monkeypatch): monkeypatch.setenv("CLANG_FORMAT_EXECUTABLE", "false") assert formatting.format_cpp_source(UNFORMATTED_CPP) == UNFORMATTED_CPP + + +@pytest.mark.skipif(shutil.which("false") is None, reason="`false` executable not available") +def test_format_cpp_source_formatting_failure(monkeypatch): + # The executable passes the availability probe but fails when formatting. + monkeypatch.setattr(formatting, "_get_clang_format", lambda: "false") + assert formatting.format_cpp_source(UNFORMATTED_CPP) == UNFORMATTED_CPP + + +def test_format_cpp_source_probes_executable_once(monkeypatch): + monkeypatch.setenv("CLANG_FORMAT_EXECUTABLE", "/nonexistent") + formatting.format_cpp_source(UNFORMATTED_CPP) + formatting.format_cpp_source(UNFORMATTED_CPP) + assert formatting._get_clang_format.cache_info().misses == 1