diff --git a/src/gt4py/next/config.py b/src/gt4py/next/config.py index 56eec687a0..308193e43c 100644 --- a/src/gt4py/next/config.py +++ b/src/gt4py/next/config.py @@ -106,6 +106,13 @@ def env_flag_to_int(name: str, default: int) -> int: ) +#: Run source formatters (e.g. black, clang-format) on generated code. +#: Only affects the readability of the generated code, never its semantics. +#: The value is captured when a code spec or workflow step is created, so +#: changing it later does not affect already existing backends. +FORMAT_SOURCES: bool = env_flag_to_bool("GT4PY_FORMAT_SOURCES", default=DEBUG) + + #: Where generated code projects should be persisted. #: Only active if BUILD_CACHE_LIFETIME is set to PERSISTENT BUILD_CACHE_DIR: pathlib.Path = ( diff --git a/src/gt4py/next/otf/artifacts.py b/src/gt4py/next/otf/artifacts.py index 204e2e75d0..ba59ed1191 100644 --- a/src/gt4py/next/otf/artifacts.py +++ b/src/gt4py/next/otf/artifacts.py @@ -30,6 +30,7 @@ from typing import Any, Generic, Optional, Protocol, TypeAlias, TypeVar, runtime_checkable from gt4py.eve import codegen +from gt4py.next import config from gt4py.next.otf.binding import interface @@ -38,15 +39,19 @@ 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``. + Formatting will happen through `eve.codegen.format_source`, only if + `.format_source` is true. For available formatting options, check the + options of the specific formatter used depending on `.formatter_key`. + + `.format_source` defaults to the value of `config.FORMAT_SOURCES` at + creation time, and it is part of the spec (and thus of its fingerprint). """ source_language: str file_extension: str formatter_key: str | None = None formatter_options: Mapping[str, Any] | None = None + format_source: bool = dataclasses.field(default_factory=lambda: config.FORMAT_SOURCES) @dataclasses.dataclass(frozen=True, kw_only=True) @@ -71,6 +76,8 @@ class SDFGCodeSpec(SourceCodeSpec): source_language: str = "SDFG" file_extension: str = "sdfg" + # There is no SDFG formatter: pin to `False`, independent of `config.FORMAT_SOURCES` + format_source: bool = dataclasses.field(default=False, init=False) @dataclasses.dataclass(frozen=True, kw_only=True) @@ -118,12 +125,15 @@ class HIPCodeSpec(CPPLikeCodeSpec): 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 {}) - ) + """Format `source` as configured in `source_code_spec` (no-op if `.format_source` is false).""" + if source_code_spec.format_source: + assert source_code_spec.formatter_key is not None, ( + "No formatter key specified in source code specification." + ) + source = codegen.format_source( + source_code_spec.formatter_key, source, **(source_code_spec.formatter_options or {}) + ) + return source CodeSpecT = TypeVar("CodeSpecT", bound=SourceCodeSpec) diff --git a/src/gt4py/next/otf/compilation/build_systems/compiledb.py b/src/gt4py/next/otf/compilation/build_systems/compiledb.py index 92a9facb08..a9060075e9 100644 --- a/src/gt4py/next/otf/compilation/build_systems/compiledb.py +++ b/src/gt4py/next/otf/compilation/build_systems/compiledb.py @@ -65,7 +65,9 @@ def __call__( deps=source.library_deps, build_type=self.cmake_build_type, cmake_flags=self.cmake_extra_flags or [], - code_spec=source.program_source.code_spec, + # The compiledb does not depend on source formatting: normalize it so the + # cache folder (keyed by the prototype source) is shared across settings. + code_spec=dataclasses.replace(source.program_source.code_spec, format_source=False), ) compiledb_template = _cc_get_compiledb( 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..3eba92252f 100644 --- a/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py +++ b/src/gt4py/next/program_processors/codegens/gtfn/gtfn_module.py @@ -9,13 +9,13 @@ from __future__ import annotations import dataclasses +from collections.abc import Callable from typing import Any, Final, Optional import factory 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 @@ -29,6 +29,31 @@ GENERATED_CONNECTIVITY_PARAM_PREFIX = "gt_conn_" +# Device-dependent settings of the generated code. Code specs are stored as factories +# (called per translation step) since their defaults depend on `config`. +_DEFAULT_CODE_SPEC_FACTORIES: Final[ + dict[core_defs.DeviceType, Callable[[], artifacts.HeaderAndSourceCodeSpec]] +] = { + core_defs.DeviceType.CPU: artifacts.CPPCodeSpec, + core_defs.DeviceType.CUDA: artifacts.CUDACodeSpec, + core_defs.DeviceType.ROCM: artifacts.HIPCodeSpec, +} +_BACKEND_HEADERS: Final[dict[core_defs.DeviceType, str]] = { + core_defs.DeviceType.CPU: "gridtools/fn/backend/naive.hpp", + core_defs.DeviceType.CUDA: "gridtools/fn/backend/gpu.hpp", + core_defs.DeviceType.ROCM: "gridtools/fn/backend/gpu.hpp", +} +_BACKEND_TYPES: Final[dict[core_defs.DeviceType, str]] = { + core_defs.DeviceType.CPU: "gridtools::fn::backend::naive{}", + core_defs.DeviceType.CUDA: "gridtools::fn::backend::gpu{}", + core_defs.DeviceType.ROCM: "gridtools::fn::backend::gpu{}", +} +_LIBRARY_NAMES: Final[dict[core_defs.DeviceType, str]] = { + core_defs.DeviceType.CPU: "gridtools_cpu", + core_defs.DeviceType.CUDA: "gridtools_gpu", + core_defs.DeviceType.ROCM: "gridtools_gpu", +} + def get_param_description(name: str, type_: Any) -> interface.Parameter: return interface.Parameter(name, type_) @@ -52,16 +77,23 @@ class GTFNTranslationStep( symbolic_domain_sizes: dict[str, itir.Expr] | None = None use_max_domain_range_on_unstructured_shift: bool | None = None - def _default_code_spec(self) -> artifacts.HeaderAndSourceCodeSpec: - match self.device_type: - case core_defs.DeviceType.CUDA: - return artifacts.CUDACodeSpec() - case core_defs.DeviceType.ROCM: - return artifacts.HIPCodeSpec() - case core_defs.DeviceType.CPU: - return artifacts.CPPCodeSpec() - case _: - raise self._not_implemented_for_device_type() + def __post_init__(self) -> None: + if (code_spec_factory := _DEFAULT_CODE_SPEC_FACTORIES.get(self.device_type)) is None: + raise NotImplementedError( + f"{self.__class__.__name__} is not implemented for device type " + f"{self.device_type.name}" + ) + # Resolve the default code spec eagerly, so its settings (e.g. `format_source`, + # which follows `config.FORMAT_SOURCES`) are part of the step and its fingerprint. + default_code_spec = code_spec_factory() + if self.code_spec is None: + object.__setattr__(self, "code_spec", default_code_spec) + elif not isinstance(self.code_spec, type(default_code_spec)): + raise ValueError( + f"Code spec '{type(self.code_spec).__name__}' does not match device type " + f"'{self.device_type.name}' (expected '{type(default_code_spec).__name__}'). " + "When replacing the device type, pass 'code_spec=None' to use the default spec." + ) def _process_regular_arguments( self, @@ -171,14 +203,15 @@ 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 ) -> artifacts.ProgramSource[artifacts.HeaderAndSourceCodeSpec]: """Generate GTFN C++ code from the ITIR definition.""" program: itir.Program = inp.data + code_spec = self.code_spec + assert code_spec is not None # resolved in `__post_init__` # handle regular parameters and arguments of the program (i.e. what the user defined in # the program) @@ -195,7 +228,7 @@ def __call__( # combine into a format that is aligned with what the backend expects parameters: list[interface.Parameter] = regular_parameters + connectivity_parameters - backend_arg = self._backend_type() + backend_arg = _BACKEND_TYPES[self.device_type] args_expr: list[str] = [backend_arg, *regular_args_expr] function = interface.Function(program.id, tuple(parameters)) @@ -210,9 +243,9 @@ def __call__( inp.args.column_axis, ) source_code = artifacts.format_source( - self._code_spec(), + code_spec, f""" - #include <{self._backend_header()}> + #include <{_BACKEND_HEADERS[self.device_type]}> #include {stencil_src} {decl_src} @@ -222,48 +255,15 @@ def __call__( module: artifacts.ProgramSource[artifacts.HeaderAndSourceCodeSpec] = ( artifacts.ProgramSource( entry_point=function, - library_deps=(interface.LibraryDependency(self._library_name(), "master"),), + library_deps=( + interface.LibraryDependency(_LIBRARY_NAMES[self.device_type], "master"), + ), source_code=source_code, - code_spec=self._code_spec(), + code_spec=code_spec, ) ) return module - def _backend_header(self) -> str: - match self.device_type: - case core_defs.DeviceType.CUDA | core_defs.DeviceType.ROCM: - return "gridtools/fn/backend/gpu.hpp" - case core_defs.DeviceType.CPU: - return "gridtools/fn/backend/naive.hpp" - case _: - raise self._not_implemented_for_device_type() - - def _backend_type(self) -> str: - match self.device_type: - case core_defs.DeviceType.CUDA | core_defs.DeviceType.ROCM: - return "gridtools::fn::backend::gpu{}" - case core_defs.DeviceType.CPU: - return "gridtools::fn::backend::naive{}" - case _: - raise self._not_implemented_for_device_type() - - def _code_spec(self) -> artifacts.HeaderAndSourceCodeSpec: - return self.code_spec if self.code_spec is not None else self._default_code_spec() - - def _library_name(self) -> str: - match self.device_type: - case core_defs.DeviceType.CUDA | core_defs.DeviceType.ROCM: - return "gridtools_gpu" - case core_defs.DeviceType.CPU: - return "gridtools_cpu" - case _: - raise self._not_implemented_for_device_type() - - def _not_implemented_for_device_type(self) -> NotImplementedError: - return NotImplementedError( - f"{self.__class__.__name__} is not implemented for device type {self.device_type.name}" - ) - class GTFNTranslationStepFactory(factory.Factory[GTFNTranslationStep]): class Meta: diff --git a/src/gt4py/next/program_processors/formatters/gtfn.py b/src/gt4py/next/program_processors/formatters/gtfn.py index 75494a1759..2932a742b7 100644 --- a/src/gt4py/next/program_processors/formatters/gtfn.py +++ b/src/gt4py/next/program_processors/formatters/gtfn.py @@ -8,6 +8,7 @@ from typing import Any +from gt4py.eve import codegen from gt4py.next.iterator import ir as itir from gt4py.next.program_processors import program_formatter from gt4py.next.program_processors.codegens.gtfn.gtfn_module import GTFNTranslationStep @@ -18,8 +19,10 @@ def format_cpp(program: itir.Program, *args: Any, **kwargs: Any) -> str: gtfn_translation = gtfn.GTFNCompileWorkflowFactory(cached_translation=False).translation assert isinstance(gtfn_translation, GTFNTranslationStep) - return gtfn_translation.generate_stencil_source( + generated_code = gtfn_translation.generate_stencil_source( program, offset_provider=kwargs.get("offset_provider", {}), column_axis=kwargs.get("column_axis", None), ) + # The purpose of this formatter is producing human-readable code, so always format it + return codegen.format_source("cpp", generated_code, style="LLVM") diff --git a/src/gt4py/next/program_processors/runners/roundtrip.py b/src/gt4py/next/program_processors/runners/roundtrip.py index 09f173d3f9..9e1eeb42ac 100644 --- a/src/gt4py/next/program_processors/runners/roundtrip.py +++ b/src/gt4py/next/program_processors/runners/roundtrip.py @@ -110,13 +110,15 @@ 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 mode (debug modules are loaded +# from a temporary file) so re-exec is skipped within a process. +_MODULE_CACHE: dict[tuple[str, bool], types.ModuleType] = {} def _generate_source( ir: itir.Program, debug: bool, + format_source: bool, use_embedded: bool, offset_provider: common.OffsetProvider, transforms: itir_transforms.GTIRTransform, @@ -128,7 +130,7 @@ def _generate_source( ( ir, transforms, - debug, + format_source, use_embedded, tuple(common.offset_provider_to_type(offset_provider).items()), ) @@ -142,9 +144,8 @@ 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: + # format output for better debuggability (e.g. line numbers, overview in the debugger). + if format_source: program = codegen.format_python_source(program) offset_literals: Iterable[str] = ( @@ -187,8 +188,8 @@ def _generate_source( def _load_module(source_code: str, debug: bool) -> types.ModuleType: - if source_code in _MODULE_CACHE: - return _MODULE_CACHE[source_code] + if (cache_key := (source_code, debug)) in _MODULE_CACHE: + return _MODULE_CACHE[cache_key] if debug: # Write to a real .py so debuggers/tracebacks have file/line info. @@ -205,7 +206,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 @@ -258,6 +259,7 @@ class Roundtrip(workflow.Workflow[stages.CompilableProgramDef, RoundtripArtifact use_embedded: bool = True dispatch_backend: Optional[next_backend.Backend] = None transforms: itir_transforms.GTIRTransform = itir_transforms.apply_common_transforms # type: ignore[assignment] # TODO(havogt): cleanup interface of `apply_common_transforms` + format_source: bool = dataclasses.field(default_factory=lambda: config.FORMAT_SOURCES) def __call__(self, inp: stages.CompilableProgramDef) -> RoundtripArtifact: debug = config.DEBUG if self.debug is None else self.debug @@ -266,6 +268,7 @@ def __call__(self, inp: stages.CompilableProgramDef) -> RoundtripArtifact: inp.data, offset_provider=inp.args.offset_provider, debug=debug, + format_source=self.format_source, use_embedded=self.use_embedded, transforms=self.transforms, ) diff --git a/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/test_compiledb.py b/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/test_compiledb.py index b208b7cd6f..9edc6c0a59 100644 --- a/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/test_compiledb.py +++ b/tests/next_tests/unit_tests/otf_tests/compilation_tests/build_systems_tests/test_compiledb.py @@ -6,6 +6,7 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import dataclasses import pathlib import shutil import tempfile @@ -91,3 +92,28 @@ def test_compiledb_project_is_relocatable(extension_source_example, clean_compil assert hasattr( importer.import_from_path(relocated_dir / new_data.module), new_data.entry_point_name ) + + +def test_compiledb_prototype_ignores_format_source(monkeypatch, extension_source_example): + prototypes = [] + + def fake_get_compiledb(renew_compiledb, prototype_program_source, **kwargs): + prototypes.append(prototype_program_source) + return pathlib.Path("compile_commands.json") + + monkeypatch.setattr(compiledb, "_cc_get_compiledb", fake_get_compiledb) + + program_source = extension_source_example.program_source + for format_source in (True, False): + code_spec = dataclasses.replace(program_source.code_spec, format_source=format_source) + compiledb.CompiledbFactory()( + dataclasses.replace( + extension_source_example, + program_source=dataclasses.replace(program_source, code_spec=code_spec), + ), + cache_lifetime=config.BuildCacheLifetime.SESSION, + ) + + assert fingerprinting.strict_fingerprinter( + prototypes[0] + ) == fingerprinting.strict_fingerprinter(prototypes[1]) diff --git a/tests/next_tests/unit_tests/otf_tests/test_languages.py b/tests/next_tests/unit_tests/otf_tests/test_languages.py index 08a738d017..7c69c73805 100644 --- a/tests/next_tests/unit_tests/otf_tests/test_languages.py +++ b/tests/next_tests/unit_tests/otf_tests/test_languages.py @@ -8,6 +8,7 @@ import pytest +from gt4py.next import config, fingerprinting from gt4py.next.otf import artifacts from gt4py.next.otf.binding import interface @@ -19,3 +20,43 @@ def test_header_files_settings_with_cpp_accepted(): library_deps=(), code_spec=artifacts.CPPCodeSpec(), ) + + +@pytest.mark.parametrize("flag", [True, False]) +def test_format_source_default_captures_config_at_creation(monkeypatch, flag): + monkeypatch.setattr(config, "FORMAT_SOURCES", flag) + code_spec = artifacts.CPPCodeSpec() + assert code_spec.format_source is flag + + # changing the global option afterwards does not affect existing specs + monkeypatch.setattr(config, "FORMAT_SOURCES", not flag) + assert code_spec.format_source is flag + assert artifacts.CPPCodeSpec().format_source is not flag + + +@pytest.mark.parametrize("flag", [True, False]) +def test_format_source_is_part_of_fingerprint(monkeypatch, flag): + monkeypatch.setattr(config, "FORMAT_SOURCES", flag) + assert fingerprinting.strict_fingerprinter( + artifacts.CPPCodeSpec() + ) != fingerprinting.strict_fingerprinter(artifacts.CPPCodeSpec(format_source=not flag)) + + +@pytest.mark.parametrize("flag", [True, False]) +def test_sdfg_code_spec_ignores_format_sources_config(monkeypatch, flag): + monkeypatch.setattr(config, "FORMAT_SOURCES", flag) + assert artifacts.SDFGCodeSpec().format_source is False + + +def test_sdfg_code_spec_rejects_format_source(): + with pytest.raises(TypeError): + artifacts.SDFGCodeSpec(format_source=True) + + +def test_format_source_follows_code_spec(): + source = "x=( 1 )" + + assert artifacts.format_source(artifacts.PythonCodeSpec(format_source=False), source) == source + assert ( + artifacts.format_source(artifacts.PythonCodeSpec(format_source=True), source) == "x = 1\n" + ) diff --git a/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py b/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py index fc077b3a90..5ac71b3af2 100644 --- a/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py +++ b/tests/next_tests/unit_tests/program_processor_tests/codegens_tests/gtfn_tests/test_gtfn_module.py @@ -13,11 +13,13 @@ import pytest import gt4py.next as gtx +from gt4py.eve import codegen from gt4py.next.iterator import builtins, ir as itir from gt4py.next.iterator.ir_utils import ir_makers as im -from gt4py.next import fingerprinting +from gt4py.next import config, fingerprinting from gt4py.next.otf import arguments, artifacts, stages from gt4py.next.program_processors.codegens.gtfn import gtfn_module +from gt4py.next.program_processors.formatters import gtfn as gtfn_formatters from gt4py.next.program_processors.runners import gtfn from gt4py.next.type_system import type_translation from gt4py.next import custom_layout_allocators as next_allocators @@ -86,6 +88,44 @@ def test_codegen(program_example): assert isinstance(module.code_spec, artifacts.CPPCodeSpec) +@pytest.mark.parametrize("flag", [True, False]) +def test_code_spec_is_resolved_at_construction(monkeypatch, flag): + monkeypatch.setattr(config, "FORMAT_SOURCES", flag) + step = gtfn_module.GTFNTranslationStep() + monkeypatch.setattr(config, "FORMAT_SOURCES", not flag) + + assert isinstance(step.code_spec, artifacts.CPPCodeSpec) + assert step.code_spec.format_source is flag + # the formatting setting is part of the step fingerprint, and thus of translation cache keys + assert fingerprinting.strict_fingerprinter(step) != fingerprinting.strict_fingerprinter( + gtfn_module.GTFNTranslationStep() + ) + + +@pytest.mark.parametrize("flag", [True, False]) +def test_format_cpp_always_formats_once(monkeypatch, program_example, flag): + monkeypatch.setattr(config, "FORMAT_SOURCES", flag) + formatted = [] + + def spy_format_source(language, source, **kwargs): + formatted.append((language, kwargs)) + return source + + monkeypatch.setattr(codegen, "format_source", spy_format_source) + + program, _ = program_example + gtfn_formatters.format_cpp(program, offset_provider={}) + + assert formatted == [("cpp", {"style": "LLVM"})] + + +def test_code_spec_device_type_mismatch(): + with pytest.raises(ValueError, match="does not match device type"): + gtfn_module.GTFNTranslationStep( + device_type=gtx.DeviceType.CUDA, code_spec=artifacts.CPPCodeSpec() + ) + + def test_hash_and_diskcache(program_example, tmp_path): fencil, parameters = program_example compilable_program = stages.CompilableProgramDef( 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..2f0b15ff29 --- /dev/null +++ b/tests/next_tests/unit_tests/program_processor_tests/runners_tests/test_roundtrip.py @@ -0,0 +1,85 @@ +# 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.eve import codegen +from gt4py.next import config +from gt4py.next.iterator import ir as itir +from gt4py.next.otf import arguments, stages +from gt4py.next.program_processors.runners import roundtrip + + +@pytest.fixture +def empty_program(): + return itir.Program(id="empty", function_definitions=[], params=[], declarations=[], body=[]) + + +@pytest.fixture +def formatter_spy(monkeypatch): + formatted_sources = [] + + def spy_formatter(source, **kwargs): + formatted_sources.append(source) + return source + + monkeypatch.setattr(codegen, "format_python_source", spy_formatter) + monkeypatch.setattr(roundtrip, "_SOURCE_CACHE", {}) + return formatted_sources + + +@pytest.mark.parametrize("flag", [True, False]) +def test_format_source_default_captures_config_at_creation(monkeypatch, flag): + monkeypatch.setattr(config, "FORMAT_SOURCES", flag) + step = roundtrip.Roundtrip() + monkeypatch.setattr(config, "FORMAT_SOURCES", not flag) + + assert step.format_source is flag + + +@pytest.mark.parametrize("format_source", [True, False]) +def test_generate_source_formatting(empty_program, formatter_spy, format_source): + roundtrip._generate_source( + empty_program, + debug=False, + format_source=format_source, + use_embedded=True, + offset_provider={}, + transforms=lambda ir, offset_provider: ir, + ) + + assert len(formatter_spy) == int(format_source) + + +@pytest.mark.parametrize("format_source", [True, False]) +def test_roundtrip_step_formatting(empty_program, formatter_spy, format_source): + step = roundtrip.Roundtrip( + transforms=lambda ir, offset_provider: ir, format_source=format_source + ) + step( + stages.CompilableProgramDef( + data=empty_program, + args=arguments.CompileTimeArgs.from_concrete(offset_provider={}), + ) + ) + + assert len(formatter_spy) == int(format_source) + + +def test_load_module_cache_respects_debug_mode(monkeypatch, tmp_path): + monkeypatch.setattr(roundtrip, "_MODULE_CACHE", {}) + monkeypatch.setattr(roundtrip.tempfile, "tempdir", str(tmp_path)) + source_code = "x = 1\n" + + roundtrip._load_module(source_code, debug=False) + assert not list(tmp_path.glob("*.py")) + + # loading the same source in debug mode must still write the temporary `.py` file + debug_module = roundtrip._load_module(source_code, debug=True) + assert list(tmp_path.glob("*.py")) + assert debug_module.__file__.startswith(str(tmp_path)) diff --git a/tests/next_tests/unit_tests/test_config.py b/tests/next_tests/unit_tests/test_config.py index a33bd5734a..2f23a6fdf3 100644 --- a/tests/next_tests/unit_tests/test_config.py +++ b/tests/next_tests/unit_tests/test_config.py @@ -6,6 +6,7 @@ # Please, refer to the LICENSE file in the root directory. # SPDX-License-Identifier: BSD-3-Clause +import importlib.util import os import pytest @@ -46,3 +47,32 @@ def test_env_flag_to_bool_invalid(env_var): def test_env_flag_to_bool_unset(env_var): _ = os.environ.pop(env_var, None) assert config.env_flag_to_bool(env_var, default=False) is False + + +def _load_fresh_config(): + """Execute a fresh copy of the `config` module, leaving `gt4py.next.config` untouched.""" + spec = importlib.util.spec_from_file_location("_fresh_gt4py_next_config", config.__file__) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +@pytest.mark.parametrize( + "debug, format_sources, expected", + [ + (None, None, False), + ("1", None, True), + ("0", "1", True), + ("1", "0", False), + ], +) +def test_format_sources_precedence(monkeypatch, debug, format_sources, expected): + # avoid adding global warning filters when executing the module + monkeypatch.setenv("GT4PY_SKIP_DACE_WARNINGS", "0") + for name, value in (("GT4PY_DEBUG", debug), ("GT4PY_FORMAT_SOURCES", format_sources)): + if value is None: + monkeypatch.delenv(name, raising=False) + else: + monkeypatch.setenv(name, value) + + assert _load_fresh_config().FORMAT_SOURCES is expected