From 2261e98953dd3ade53febf253027c2235e864d5a Mon Sep 17 00:00:00 2001 From: Mohamed Ashraf Date: Tue, 3 Mar 2026 20:33:57 +0000 Subject: [PATCH 1/4] fix: suppress Error Prone CheckReturnValue in instrumented tests and fix stale pom dependency Add @SuppressWarnings("CheckReturnValue") to all generated instrumented test classes. Projects using Error Prone (e.g. Guava) enforce CheckReturnValue as a compiler error, which rejects our performance-only tests that intentionally discard return values after assertion stripping. Also fix add_codeflash_dependency_to_pom to detect and replace stale system-scope dependencies left by previous runs with the correct test scope. Co-Authored-By: Claude Opus 4.6 --- codeflash/languages/java/build_tools.py | 21 + codeflash/languages/java/instrumentation.py | 24 + .../test_java/test_instrumentation.py | 486 +++++++++--------- 3 files changed, 291 insertions(+), 240 deletions(-) diff --git a/codeflash/languages/java/build_tools.py b/codeflash/languages/java/build_tools.py index ba4a5ccd4..571b6c2ba 100644 --- a/codeflash/languages/java/build_tools.py +++ b/codeflash/languages/java/build_tools.py @@ -8,6 +8,7 @@ import logging import os +import re import shutil import subprocess import xml.etree.ElementTree as ET @@ -645,6 +646,26 @@ def add_codeflash_dependency_to_pom(pom_path: Path) -> bool: # Check if already present if "codeflash-runtime" in content: + # If a previous run left a system-scope dependency, replace it with test scope. + # System-scope dependencies cause Maven warnings and are rejected by some projects. + if "system" in content and "codeflash-runtime" in content: + content = re.sub( + r"\s*com\.codeflash\s*" + r"codeflash-runtime\s*" + r"[^<]*\s*" + r"system\s*" + r"[^<]*\s*", + "\n" + " com.codeflash\n" + " codeflash-runtime\n" + " 1.0.0\n" + " test\n" + " ", + content, + ) + pom_path.write_text(content, encoding="utf-8") + logger.info("Replaced system-scope codeflash-runtime dependency with test scope") + return True logger.info("codeflash-runtime dependency already present in pom.xml") return True diff --git a/codeflash/languages/java/instrumentation.py b/codeflash/languages/java/instrumentation.py index ee7700f5e..d5af758c9 100644 --- a/codeflash/languages/java/instrumentation.py +++ b/codeflash/languages/java/instrumentation.py @@ -638,6 +638,12 @@ def instrument_existing_test( # replacing substrings of other identifiers. modified_source = re.sub(rf"\b{re.escape(original_class_name)}\b", new_class_name, source) + # Add @SuppressWarnings("CheckReturnValue") to the class declaration. + # Projects using Error Prone (e.g. Guava) enforce CheckReturnValue as a compiler error. + # Our instrumented tests intentionally discard return values in performance-only mode + # (after assertion stripping), which would fail compilation without this suppression. + modified_source = _add_suppress_warnings_annotation(modified_source, new_class_name) + # Add timing instrumentation to test methods # Use original class name (without suffix) in timing markers for consistency with Python if mode == "performance": @@ -828,6 +834,21 @@ def _add_behavior_instrumentation(source: str, class_name: str, func_name: str) return "\n".join(result) +def _add_suppress_warnings_annotation(source: str, class_name: str) -> str: + """Add @SuppressWarnings("CheckReturnValue") before the class declaration. + + Projects using Error Prone (e.g. Guava) enforce CheckReturnValue as a compiler error. + Our instrumented tests intentionally discard return values after assertion stripping, + which would fail compilation without this suppression. + """ + class_decl_pattern = re.compile(rf"^((?:public\s+)?class\s+{re.escape(class_name)}\b)", re.MULTILINE) + match = class_decl_pattern.search(source) + if not match: + return source + insert_pos = match.start() + return source[:insert_pos] + '@SuppressWarnings("CheckReturnValue")\n' + source[insert_pos:] + + def _add_timing_instrumentation(source: str, class_name: str, func_name: str) -> str: """Add timing instrumentation to test methods with inner loop for JIT warmup. @@ -1307,6 +1328,9 @@ def instrument_generated_java_test( # This includes the class declaration, return types, constructor calls, etc. modified_code = re.sub(rf"\b{re.escape(original_class_name)}\b", new_class_name, test_code) + # Suppress Error Prone's CheckReturnValue for generated performance tests + modified_code = _add_suppress_warnings_annotation(modified_code, new_class_name) + modified_code = _add_timing_instrumentation( modified_code, original_class_name, # Use original name in markers, not the renamed class diff --git a/tests/test_languages/test_java/test_instrumentation.py b/tests/test_languages/test_java/test_instrumentation.py index a7e1e769f..55ea7b980 100644 --- a/tests/test_languages/test_java/test_instrumentation.py +++ b/tests/test_languages/test_java/test_instrumentation.py @@ -122,10 +122,7 @@ def test_instrument_behavior_mode_simple(self, tmp_path: Path): ) success, result = instrument_existing_test( - test_string=source, - function_to_optimize=func, - mode="behavior", - test_path=test_file, + test_string=source, function_to_optimize=func, mode="behavior", test_path=test_file ) expected = """import org.junit.jupiter.api.Test; @@ -133,6 +130,7 @@ def test_instrument_behavior_mode_simple(self, tmp_path: Path): import java.sql.DriverManager; import java.sql.PreparedStatement; +@SuppressWarnings("CheckReturnValue") public class CalculatorTest__perfinstrumented { @Test public void testAdd() { @@ -234,10 +232,7 @@ def test_instrument_behavior_mode_assert_throws_expression_lambda(self, tmp_path ) success, result = instrument_existing_test( - test_string=source, - function_to_optimize=func, - mode="behavior", - test_path=test_file, + test_string=source, function_to_optimize=func, mode="behavior", test_path=test_file ) expected = """import org.junit.jupiter.api.Test; @@ -246,6 +241,7 @@ def test_instrument_behavior_mode_assert_throws_expression_lambda(self, tmp_path import java.sql.DriverManager; import java.sql.PreparedStatement; +@SuppressWarnings("CheckReturnValue") public class FibonacciTest__perfinstrumented { @Test void testNegativeInput_ThrowsIllegalArgumentException() { @@ -362,10 +358,7 @@ def test_instrument_behavior_mode_assert_throws_block_lambda(self, tmp_path: Pat ) success, result = instrument_existing_test( - test_string=source, - function_to_optimize=func, - mode="behavior", - test_path=test_file, + test_string=source, function_to_optimize=func, mode="behavior", test_path=test_file ) expected = """import org.junit.jupiter.api.Test; @@ -374,6 +367,7 @@ def test_instrument_behavior_mode_assert_throws_block_lambda(self, tmp_path: Pat import java.sql.DriverManager; import java.sql.PreparedStatement; +@SuppressWarnings("CheckReturnValue") public class FibonacciTest__perfinstrumented { @Test void testNegativeInput_ThrowsIllegalArgumentException() { @@ -481,14 +475,12 @@ def test_instrument_performance_mode_simple(self, tmp_path: Path): ) success, result = instrument_existing_test( - test_string=source, - function_to_optimize=func, - mode="performance", - test_path=test_file, + test_string=source, function_to_optimize=func, mode="performance", test_path=test_file ) expected = """import org.junit.jupiter.api.Test; +@SuppressWarnings("CheckReturnValue") public class CalculatorTest__perfonlyinstrumented { @Test public void testAdd() { @@ -553,14 +545,12 @@ def test_instrument_performance_mode_multiple_tests(self, tmp_path: Path): ) success, result = instrument_existing_test( - test_string=source, - function_to_optimize=func, - mode="performance", - test_path=test_file, + test_string=source, function_to_optimize=func, mode="performance", test_path=test_file ) expected = """import org.junit.jupiter.api.Test; +@SuppressWarnings("CheckReturnValue") public class MathTest__perfonlyinstrumented { @Test public void testAdd() { @@ -656,16 +646,14 @@ def test_instrument_preserves_annotations(self, tmp_path: Path): ) success, result = instrument_existing_test( - test_string=source, - function_to_optimize=func, - mode="performance", - test_path=test_file, + test_string=source, function_to_optimize=func, mode="performance", test_path=test_file ) expected = """import org.junit.jupiter.api.Test; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Disabled; +@SuppressWarnings("CheckReturnValue") public class ServiceTest__perfonlyinstrumented { @Test @DisplayName("Test service call") @@ -721,11 +709,7 @@ def test_missing_file(self, tmp_path: Path): ) with pytest.raises(ValueError): - instrument_existing_test( - test_string="", - function_to_optimize=func, - mode="behavior", - ) + instrument_existing_test(test_string="", function_to_optimize=func, mode="behavior") class TestKryoSerializerUsage: @@ -1154,12 +1138,7 @@ def test_create_benchmark_different_iterations(self): language="java", ) - result = create_benchmark_test( - func, - test_setup_code="", - invocation_code="multiply(5, 3)", - iterations=5000, - ) + result = create_benchmark_test(func, test_setup_code="", invocation_code="multiply(5, 3)", iterations=5000) # Note: Empty test_setup_code still has 8-space indentation on its line expected = ( @@ -1255,11 +1234,7 @@ def test_instrument_generated_test_behavior_mode(self): language="java", ) result = instrument_generated_java_test( - test_code, - function_name="add", - qualified_name="Calculator.add", - mode="behavior", - function_to_optimize=func, + test_code, function_name="add", qualified_name="Calculator.add", mode="behavior", function_to_optimize=func ) expected = """import org.junit.jupiter.api.Test; @@ -1267,6 +1242,7 @@ def test_instrument_generated_test_behavior_mode(self): import java.sql.DriverManager; import java.sql.PreparedStatement; +@SuppressWarnings("CheckReturnValue") public class CalculatorTest__perfinstrumented { @Test public void testAdd() { @@ -1360,6 +1336,7 @@ def test_instrument_generated_test_performance_mode(self): expected = """import org.junit.jupiter.api.Test; +@SuppressWarnings("CheckReturnValue") public class GeneratedTest__perfonlyinstrumented { @Test public void testMethod() { @@ -1532,14 +1509,12 @@ def test_instrumented_code_has_balanced_braces(self, tmp_path: Path): ) success, result = instrument_existing_test( - test_string=source, - function_to_optimize=func, - mode="performance", - test_path=test_file, + test_string=source, function_to_optimize=func, mode="performance", test_path=test_file ) expected = """import org.junit.jupiter.api.Test; +@SuppressWarnings("CheckReturnValue") public class BraceTest__perfonlyinstrumented { @Test public void testOne() { @@ -1613,10 +1588,7 @@ def test_instrumented_code_preserves_imports(self, tmp_path: Path): ) success, result = instrument_existing_test( - test_string=source, - function_to_optimize=func, - mode="performance", - test_path=test_file, + test_string=source, function_to_optimize=func, mode="performance", test_path=test_file ) expected = """package com.example; @@ -1626,6 +1598,7 @@ def test_instrumented_code_preserves_imports(self, tmp_path: Path): import java.util.List; import java.util.ArrayList; +@SuppressWarnings("CheckReturnValue") public class ImportTest__perfonlyinstrumented { @Test public void testCollections() { @@ -1688,14 +1661,12 @@ def test_empty_test_method(self, tmp_path: Path): ) success, result = instrument_existing_test( - test_string=source, - function_to_optimize=func, - mode="performance", - test_path=test_file, + test_string=source, function_to_optimize=func, mode="performance", test_path=test_file ) expected = """import org.junit.jupiter.api.Test; +@SuppressWarnings("CheckReturnValue") public class EmptyTest__perfonlyinstrumented { @Test public void testEmpty() { @@ -1736,14 +1707,12 @@ def test_test_with_nested_braces(self, tmp_path: Path): ) success, result = instrument_existing_test( - test_string=source, - function_to_optimize=func, - mode="performance", - test_path=test_file, + test_string=source, function_to_optimize=func, mode="performance", test_path=test_file ) expected = """import org.junit.jupiter.api.Test; +@SuppressWarnings("CheckReturnValue") public class NestedTest__perfonlyinstrumented { @Test public void testNested() { @@ -1817,15 +1786,13 @@ class InnerTests { ) success, result = instrument_existing_test( - test_string=source, - function_to_optimize=func, - mode="performance", - test_path=test_file, + test_string=source, function_to_optimize=func, mode="performance", test_path=test_file ) expected = """import org.junit.jupiter.api.Test; import org.junit.jupiter.api.Nested; +@SuppressWarnings("CheckReturnValue") public class InnerClassTest__perfonlyinstrumented { @Test public void testOuter() { @@ -1881,22 +1848,20 @@ def test_instrument_with_cjk_in_string_literal(self, tmp_path: Path): ) success, result = instrument_existing_test( - test_string=source, - function_to_optimize=func, - mode="performance", - test_path=test_file, + test_string=source, function_to_optimize=func, mode="performance", test_path=test_file ) # The blank line between _cf_fn1 and the prefix body has 8 trailing spaces # (the indent level) — this is the f"{indent}\n" separator in the instrumentation code. expected = ( - 'import org.junit.jupiter.api.Test;\n' - 'import static org.junit.jupiter.api.Assertions.*;\n' - '\n' - 'public class Utf8Test__perfonlyinstrumented {\n' - ' @Test\n' - ' public void testWithCjk() {\n' - ' // Codeflash timing instrumentation with inner loop for JIT warmup\n' + "import org.junit.jupiter.api.Test;\n" + "import static org.junit.jupiter.api.Assertions.*;\n" + "\n" + '@SuppressWarnings("CheckReturnValue")\n' + "public class Utf8Test__perfonlyinstrumented {\n" + " @Test\n" + " public void testWithCjk() {\n" + " // Codeflash timing instrumentation with inner loop for JIT warmup\n" ' int _cf_outerLoop1 = Integer.parseInt(System.getenv("CODEFLASH_LOOP_INDEX"));\n' ' int _cf_maxInnerIterations1 = Integer.parseInt(System.getenv().getOrDefault("CODEFLASH_INNER_ITERATIONS", "10"));\n' ' int _cf_innerIterations1 = Integer.parseInt(System.getenv().getOrDefault("CODEFLASH_INNER_ITERATIONS", "10"));\n' @@ -1904,25 +1869,25 @@ def test_instrument_with_cjk_in_string_literal(self, tmp_path: Path): ' String _cf_cls1 = "Utf8Test";\n' ' String _cf_test1 = "testWithCjk";\n' ' String _cf_fn1 = "compute";\n' - ' \n' + " \n" ' String label = "\u30c6\u30b9\u30c8\u540d\u524d";\n' - ' for (int _cf_i1 = 0; _cf_i1 < _cf_innerIterations1; _cf_i1++) {\n' - ' int _cf_loopId1 = _cf_outerLoop1 * _cf_maxInnerIterations1 + _cf_i1;\n' + " for (int _cf_i1 = 0; _cf_i1 < _cf_innerIterations1; _cf_i1++) {\n" + " int _cf_loopId1 = _cf_outerLoop1 * _cf_maxInnerIterations1 + _cf_i1;\n" ' System.out.println("!$######" + _cf_mod1 + ":" + _cf_cls1 + "." + _cf_test1 + ":" + _cf_fn1 + ":" + _cf_loopId1 + ":" + "1" + "######$!");\n' - ' long _cf_end1 = -1;\n' - ' long _cf_start1 = 0;\n' - ' try {\n' - ' _cf_start1 = System.nanoTime();\n' - ' assertEquals(42, compute(21));\n' - ' _cf_end1 = System.nanoTime();\n' - ' } finally {\n' - ' long _cf_end1_finally = System.nanoTime();\n' - ' long _cf_dur1 = (_cf_end1 != -1 ? _cf_end1 : _cf_end1_finally) - _cf_start1;\n' + " long _cf_end1 = -1;\n" + " long _cf_start1 = 0;\n" + " try {\n" + " _cf_start1 = System.nanoTime();\n" + " assertEquals(42, compute(21));\n" + " _cf_end1 = System.nanoTime();\n" + " } finally {\n" + " long _cf_end1_finally = System.nanoTime();\n" + " long _cf_dur1 = (_cf_end1 != -1 ? _cf_end1 : _cf_end1_finally) - _cf_start1;\n" ' System.out.println("!######" + _cf_mod1 + ":" + _cf_cls1 + "." + _cf_test1 + ":" + _cf_fn1 + ":" + _cf_loopId1 + ":" + "1" + ":" + _cf_dur1 + "######!");\n' - ' }\n' - ' }\n' - ' }\n' - '}\n' + " }\n" + " }\n" + " }\n" + "}\n" ) assert success is True assert result == expected @@ -1955,22 +1920,20 @@ def test_instrument_with_multibyte_in_comment(self, tmp_path: Path): ) success, result = instrument_existing_test( - test_string=source, - function_to_optimize=func, - mode="performance", - test_path=test_file, + test_string=source, function_to_optimize=func, mode="performance", test_path=test_file ) assert success is True expected = ( - 'import org.junit.jupiter.api.Test;\n' - 'import static org.junit.jupiter.api.Assertions.*;\n' - '\n' - 'public class AccentTest__perfonlyinstrumented {\n' - ' @Test\n' - ' public void testWithAccent() {\n' - ' // Codeflash timing instrumentation with inner loop for JIT warmup\n' + "import org.junit.jupiter.api.Test;\n" + "import static org.junit.jupiter.api.Assertions.*;\n" + "\n" + '@SuppressWarnings("CheckReturnValue")\n' + "public class AccentTest__perfonlyinstrumented {\n" + " @Test\n" + " public void testWithAccent() {\n" + " // Codeflash timing instrumentation with inner loop for JIT warmup\n" ' int _cf_outerLoop1 = Integer.parseInt(System.getenv("CODEFLASH_LOOP_INDEX"));\n' ' int _cf_maxInnerIterations1 = Integer.parseInt(System.getenv().getOrDefault("CODEFLASH_INNER_ITERATIONS", "10"));\n' ' int _cf_innerIterations1 = Integer.parseInt(System.getenv().getOrDefault("CODEFLASH_INNER_ITERATIONS", "10"));\n' @@ -1978,34 +1941,33 @@ def test_instrument_with_multibyte_in_comment(self, tmp_path: Path): ' String _cf_cls1 = "AccentTest";\n' ' String _cf_test1 = "testWithAccent";\n' ' String _cf_fn1 = "calculate";\n' - ' \n' - ' // R\u00e9sum\u00e9 processing test with accented chars\n' + " \n" + " // R\u00e9sum\u00e9 processing test with accented chars\n" ' String name = "caf\u00e9";\n' - ' for (int _cf_i1 = 0; _cf_i1 < _cf_innerIterations1; _cf_i1++) {\n' - ' int _cf_loopId1 = _cf_outerLoop1 * _cf_maxInnerIterations1 + _cf_i1;\n' + " for (int _cf_i1 = 0; _cf_i1 < _cf_innerIterations1; _cf_i1++) {\n" + " int _cf_loopId1 = _cf_outerLoop1 * _cf_maxInnerIterations1 + _cf_i1;\n" ' System.out.println("!$######" + _cf_mod1 + ":" + _cf_cls1 + "." + _cf_test1 + ":" + _cf_fn1 + ":" + _cf_loopId1 + ":" + "1" + "######$!");\n' - ' long _cf_end1 = -1;\n' - ' long _cf_start1 = 0;\n' - ' try {\n' - ' _cf_start1 = System.nanoTime();\n' - ' assertEquals(10, calculate(5));\n' - ' _cf_end1 = System.nanoTime();\n' - ' } finally {\n' - ' long _cf_end1_finally = System.nanoTime();\n' - ' long _cf_dur1 = (_cf_end1 != -1 ? _cf_end1 : _cf_end1_finally) - _cf_start1;\n' + " long _cf_end1 = -1;\n" + " long _cf_start1 = 0;\n" + " try {\n" + " _cf_start1 = System.nanoTime();\n" + " assertEquals(10, calculate(5));\n" + " _cf_end1 = System.nanoTime();\n" + " } finally {\n" + " long _cf_end1_finally = System.nanoTime();\n" + " long _cf_dur1 = (_cf_end1 != -1 ? _cf_end1 : _cf_end1_finally) - _cf_start1;\n" ' System.out.println("!######" + _cf_mod1 + ":" + _cf_cls1 + "." + _cf_test1 + ":" + _cf_fn1 + ":" + _cf_loopId1 + ":" + "1" + ":" + _cf_dur1 + "######!");\n' - ' }\n' - ' }\n' - ' }\n' - '}\n' + " }\n" + " }\n" + " }\n" + "}\n" ) assert result == expected # Skip all E2E tests if Maven is not available requires_maven = pytest.mark.skipif( - find_maven_executable() is None, - reason="Maven not found - skipping execution tests", + find_maven_executable() is None, reason="Maven not found - skipping execution tests" ) @@ -2080,6 +2042,7 @@ def java_project(self, tmp_path: Path): """Create a temporary Maven project and set up Java language context.""" # Force set the language to Java (reset the singleton first) import codeflash.languages.current as current_module + current_module._current_language = None set_current_language(Language.JAVA) @@ -2107,14 +2070,17 @@ def test_run_and_parse_behavior_mode(self, java_project): project_root, src_dir, test_dir = java_project # Create source file - (src_dir / "Calculator.java").write_text("""package com.example; + (src_dir / "Calculator.java").write_text( + """package com.example; public class Calculator { public int add(int a, int b) { return a + b; } } -""", encoding="utf-8") +""", + encoding="utf-8", + ) # Create and instrument test test_source = """package com.example; @@ -2153,32 +2119,33 @@ def test_run_and_parse_behavior_mode(self, java_project): # Create Optimizer and FunctionOptimizer fto = FunctionToOptimize( - function_name="add", - file_path=src_dir / "Calculator.java", - parents=[], - language="java", + function_name="add", file_path=src_dir / "Calculator.java", parents=[], language="java" ) - opt = Optimizer(Namespace( - project_root=project_root, - disable_telemetry=True, - tests_root=test_dir, - test_project_root=project_root, - pytest_cmd="pytest", - experiment_id=None, - )) + opt = Optimizer( + Namespace( + project_root=project_root, + disable_telemetry=True, + tests_root=test_dir, + test_project_root=project_root, + pytest_cmd="pytest", + experiment_id=None, + ) + ) func_optimizer = opt.create_function_optimizer(fto) assert func_optimizer is not None - func_optimizer.test_files = TestFiles(test_files=[ - TestFile( - instrumented_behavior_file_path=instrumented_file, - test_type=TestType.EXISTING_UNIT_TEST, - original_file_path=test_file, - benchmarking_file_path=instrumented_file, # Use same file for behavior tests - ) - ]) + func_optimizer.test_files = TestFiles( + test_files=[ + TestFile( + instrumented_behavior_file_path=instrumented_file, + test_type=TestType.EXISTING_UNIT_TEST, + original_file_path=test_file, + benchmarking_file_path=instrumented_file, # Use same file for behavior tests + ) + ] + ) # Run and parse tests test_env = os.environ.copy() @@ -2219,14 +2186,17 @@ def test_run_and_parse_performance_mode(self, java_project): project_root, src_dir, test_dir = java_project # Create source file - (src_dir / "MathUtils.java").write_text("""package com.example; + (src_dir / "MathUtils.java").write_text( + """package com.example; public class MathUtils { public int multiply(int a, int b) { return a * b; } } -""", encoding="utf-8") +""", + encoding="utf-8", + ) # Create and instrument test test_source = """package com.example; @@ -2265,6 +2235,7 @@ def test_run_and_parse_performance_mode(self, java_project): import org.junit.jupiter.api.Test; import static org.junit.jupiter.api.Assertions.*; +@SuppressWarnings("CheckReturnValue") public class MathUtilsTest__perfonlyinstrumented { @Test public void testMultiply() { @@ -2303,32 +2274,33 @@ def test_run_and_parse_performance_mode(self, java_project): # Create Optimizer and FunctionOptimizer fto = FunctionToOptimize( - function_name="multiply", - file_path=src_dir / "MathUtils.java", - parents=[], - language="java", + function_name="multiply", file_path=src_dir / "MathUtils.java", parents=[], language="java" ) - opt = Optimizer(Namespace( - project_root=project_root, - disable_telemetry=True, - tests_root=test_dir, - test_project_root=project_root, - pytest_cmd="pytest", - experiment_id=None, - )) + opt = Optimizer( + Namespace( + project_root=project_root, + disable_telemetry=True, + tests_root=test_dir, + test_project_root=project_root, + pytest_cmd="pytest", + experiment_id=None, + ) + ) func_optimizer = opt.create_function_optimizer(fto) assert func_optimizer is not None - func_optimizer.test_files = TestFiles(test_files=[ - TestFile( - instrumented_behavior_file_path=test_file, - test_type=TestType.EXISTING_UNIT_TEST, - original_file_path=test_file, - benchmarking_file_path=instrumented_file, - ) - ]) + func_optimizer.test_files = TestFiles( + test_files=[ + TestFile( + instrumented_behavior_file_path=test_file, + test_type=TestType.EXISTING_UNIT_TEST, + original_file_path=test_file, + benchmarking_file_path=instrumented_file, + ) + ] + ) # Run performance tests with inner_iterations=2 for fast test test_env = os.environ.copy() @@ -2377,14 +2349,17 @@ def test_run_and_parse_multiple_test_methods(self, java_project): project_root, src_dir, test_dir = java_project # Create source file - (src_dir / "StringUtils.java").write_text("""package com.example; + (src_dir / "StringUtils.java").write_text( + """package com.example; public class StringUtils { public String reverse(String s) { return new StringBuilder(s).reverse().toString(); } } -""", encoding="utf-8") +""", + encoding="utf-8", + ) # Create test with multiple methods test_source = """package com.example; @@ -2431,30 +2406,31 @@ def test_run_and_parse_multiple_test_methods(self, java_project): instrumented_file.write_text(instrumented, encoding="utf-8") fto = FunctionToOptimize( - function_name="reverse", - file_path=src_dir / "StringUtils.java", - parents=[], - language="java", + function_name="reverse", file_path=src_dir / "StringUtils.java", parents=[], language="java" ) - opt = Optimizer(Namespace( - project_root=project_root, - disable_telemetry=True, - tests_root=test_dir, - test_project_root=project_root, - pytest_cmd="pytest", - experiment_id=None, - )) + opt = Optimizer( + Namespace( + project_root=project_root, + disable_telemetry=True, + tests_root=test_dir, + test_project_root=project_root, + pytest_cmd="pytest", + experiment_id=None, + ) + ) func_optimizer = opt.create_function_optimizer(fto) - func_optimizer.test_files = TestFiles(test_files=[ - TestFile( - instrumented_behavior_file_path=instrumented_file, - test_type=TestType.EXISTING_UNIT_TEST, - original_file_path=test_file, - benchmarking_file_path=instrumented_file, # Use same file for behavior tests - ) - ]) + func_optimizer.test_files = TestFiles( + test_files=[ + TestFile( + instrumented_behavior_file_path=instrumented_file, + test_type=TestType.EXISTING_UNIT_TEST, + original_file_path=test_file, + benchmarking_file_path=instrumented_file, # Use same file for behavior tests + ) + ] + ) test_env = os.environ.copy() test_env["CODEFLASH_TEST_ITERATION"] = "0" @@ -2488,14 +2464,17 @@ def test_run_and_parse_failing_test(self, java_project): project_root, src_dir, test_dir = java_project # Create source file with a bug - (src_dir / "BrokenCalc.java").write_text("""package com.example; + (src_dir / "BrokenCalc.java").write_text( + """package com.example; public class BrokenCalc { public int add(int a, int b) { return a + b + 1; // Bug: adds extra 1 } } -""", encoding="utf-8") +""", + encoding="utf-8", + ) # Create test that will fail test_source = """package com.example; @@ -2533,30 +2512,31 @@ def test_run_and_parse_failing_test(self, java_project): instrumented_file.write_text(instrumented, encoding="utf-8") fto = FunctionToOptimize( - function_name="add", - file_path=src_dir / "BrokenCalc.java", - parents=[], - language="java", + function_name="add", file_path=src_dir / "BrokenCalc.java", parents=[], language="java" ) - opt = Optimizer(Namespace( - project_root=project_root, - disable_telemetry=True, - tests_root=test_dir, - test_project_root=project_root, - pytest_cmd="pytest", - experiment_id=None, - )) + opt = Optimizer( + Namespace( + project_root=project_root, + disable_telemetry=True, + tests_root=test_dir, + test_project_root=project_root, + pytest_cmd="pytest", + experiment_id=None, + ) + ) func_optimizer = opt.create_function_optimizer(fto) - func_optimizer.test_files = TestFiles(test_files=[ - TestFile( - instrumented_behavior_file_path=instrumented_file, - test_type=TestType.EXISTING_UNIT_TEST, - original_file_path=test_file, - benchmarking_file_path=instrumented_file, # Use same file for behavior tests - ) - ]) + func_optimizer.test_files = TestFiles( + test_files=[ + TestFile( + instrumented_behavior_file_path=instrumented_file, + test_type=TestType.EXISTING_UNIT_TEST, + original_file_path=test_file, + benchmarking_file_path=instrumented_file, # Use same file for behavior tests + ) + ] + ) test_env = os.environ.copy() test_env["CODEFLASH_TEST_ITERATION"] = "0" @@ -2594,7 +2574,8 @@ def test_behavior_mode_writes_to_sqlite(self, java_project): project_root, src_dir, test_dir = java_project # Create source file - (src_dir / "Counter.java").write_text("""package com.example; + (src_dir / "Counter.java").write_text( + """package com.example; public class Counter { private int value = 0; @@ -2603,7 +2584,9 @@ def test_behavior_mode_writes_to_sqlite(self, java_project): return ++value; } } -""", encoding="utf-8") +""", + encoding="utf-8", + ) # Create test file - single test method for simplicity test_source = """package com.example; @@ -2646,6 +2629,7 @@ def test_behavior_mode_writes_to_sqlite(self, java_project): import java.sql.DriverManager; import java.sql.PreparedStatement; +@SuppressWarnings("CheckReturnValue") public class CounterTest__perfinstrumented { @Test public void testIncrement() { @@ -2715,32 +2699,33 @@ def test_behavior_mode_writes_to_sqlite(self, java_project): # Create Optimizer and FunctionOptimizer fto = FunctionToOptimize( - function_name="increment", - file_path=src_dir / "Counter.java", - parents=[], - language="java", + function_name="increment", file_path=src_dir / "Counter.java", parents=[], language="java" ) - opt = Optimizer(Namespace( - project_root=project_root, - disable_telemetry=True, - tests_root=test_dir, - test_project_root=project_root, - pytest_cmd="pytest", - experiment_id=None, - )) + opt = Optimizer( + Namespace( + project_root=project_root, + disable_telemetry=True, + tests_root=test_dir, + test_project_root=project_root, + pytest_cmd="pytest", + experiment_id=None, + ) + ) func_optimizer = opt.create_function_optimizer(fto) assert func_optimizer is not None - func_optimizer.test_files = TestFiles(test_files=[ - TestFile( - instrumented_behavior_file_path=instrumented_file, - test_type=TestType.EXISTING_UNIT_TEST, - original_file_path=test_file, - benchmarking_file_path=instrumented_file, - ) - ]) + func_optimizer.test_files = TestFiles( + test_files=[ + TestFile( + instrumented_behavior_file_path=instrumented_file, + test_type=TestType.EXISTING_UNIT_TEST, + original_file_path=test_file, + benchmarking_file_path=instrumented_file, + ) + ] + ) # Run tests test_env = os.environ.copy() @@ -2766,11 +2751,13 @@ def test_behavior_mode_writes_to_sqlite(self, java_project): # Find the SQLite file that was created # SQLite is created at get_run_tmp_file path from codeflash.code_utils.code_utils import get_run_tmp_file + sqlite_file = get_run_tmp_file(Path("test_return_values_0.sqlite")) if not sqlite_file.exists(): # Fall back to checking temp directory for any SQLite files import tempfile + sqlite_files = list(Path(tempfile.gettempdir()).glob("**/test_return_values_*.sqlite")) assert len(sqlite_files) >= 1, f"SQLite file should have been created at {sqlite_file} or in temp dir" sqlite_file = max(sqlite_files, key=lambda p: p.stat().st_mtime) @@ -2789,8 +2776,17 @@ def test_behavior_mode_writes_to_sqlite(self, java_project): rows = cursor.fetchall() for row in rows: - test_module_path, test_class_name, test_function_name, function_getting_tested, \ - loop_index, iteration_id, runtime, return_value, verification_type = row + ( + test_module_path, + test_class_name, + test_function_name, + function_getting_tested, + loop_index, + iteration_id, + runtime, + return_value, + verification_type, + ) = row # Verify fields assert test_module_path == "CounterTest" @@ -2819,7 +2815,8 @@ def test_performance_mode_inner_loop_timing_markers(self, java_project): project_root, src_dir, test_dir = java_project # Create a simple function to optimize - (src_dir / "Fibonacci.java").write_text("""package com.example; + (src_dir / "Fibonacci.java").write_text( + """package com.example; public class Fibonacci { public int fib(int n) { @@ -2827,7 +2824,9 @@ def test_performance_mode_inner_loop_timing_markers(self, java_project): return fib(n - 1) + fib(n - 2); } } -""", encoding="utf-8") +""", + encoding="utf-8", + ) # Create test file test_source = """package com.example; @@ -2867,6 +2866,7 @@ def test_performance_mode_inner_loop_timing_markers(self, java_project): import org.junit.jupiter.api.Test; import static org.junit.jupiter.api.Assertions.*; +@SuppressWarnings("CheckReturnValue") public class FibonacciTest__perfonlyinstrumented { @Test public void testFib() { @@ -2947,7 +2947,9 @@ def __init__(self, path): # Verify invocation IDs are constant (wrapper ID) across all inner iterations invocation_ids = [m[4] for m in start_matches] - assert all(id == invocation_ids[0] for id in invocation_ids), f"Expected constant invocation IDs, got: {invocation_ids}" + assert all(id == invocation_ids[0] for id in invocation_ids), ( + f"Expected constant invocation IDs, got: {invocation_ids}" + ) # Verify loop IDs are 2 and 3 (outerLoop=1, maxInner=2, inner=0,1 → 1*2+0=2, 1*2+1=3) loop_ids = [m[3] for m in start_matches] @@ -2968,14 +2970,17 @@ def test_performance_mode_multiple_methods_inner_loop(self, java_project): project_root, src_dir, test_dir = java_project # Create a simple math class - (src_dir / "MathOps.java").write_text("""package com.example; + (src_dir / "MathOps.java").write_text( + """package com.example; public class MathOps { public int add(int a, int b) { return a + b; } } -""", encoding="utf-8") +""", + encoding="utf-8", + ) # Create test with multiple test methods test_source = """package com.example; @@ -3079,7 +3084,8 @@ def test_time_correction_instrumentation(self, java_project): project_root, src_dir, test_dir = java_project # Create SpinWait class — Java equivalent of Python's accurate_sleepfunc - (src_dir / "SpinWait.java").write_text("""package com.example; + (src_dir / "SpinWait.java").write_text( + """package com.example; public class SpinWait { public static long spinWait(long durationNs) { @@ -3089,7 +3095,9 @@ def test_time_correction_instrumentation(self, java_project): return durationNs; } } -""", encoding="utf-8") +""", + encoding="utf-8", + ) # Two test methods with known durations — mirrors Python's parametrize with # (0.01, 0.010) and (0.02, 0.020) which map to 100ms and 200ms @@ -3125,10 +3133,7 @@ def test_time_correction_instrumentation(self, java_project): # Instrument for performance mode success, instrumented = instrument_existing_test( - test_string=test_source, - function_to_optimize=func_info, - mode="performance", - test_path=test_file, + test_string=test_source, function_to_optimize=func_info, mode="performance", test_path=test_file ) assert success, "Instrumentation should succeed" @@ -3139,6 +3144,7 @@ def test_time_correction_instrumentation(self, java_project): import org.junit.jupiter.api.Test; import static org.junit.jupiter.api.Assertions.*; +@SuppressWarnings("CheckReturnValue") public class SpinWaitTest__perfonlyinstrumented { @Test public void testSpinShort() { From b9135c5f6a02dc7b7e1acc6ae8502ec1a6549335 Mon Sep 17 00:00:00 2001 From: Mohamed Ashraf Date: Wed, 4 Mar 2026 00:00:14 +0000 Subject: [PATCH 2/4] fix: address PR review feedback for Error Prone and build_tools - Remove redundant condition check in add_codeflash_dependency_to_pom - Use lookahead-based regex to handle arbitrary XML element order in system-scope dependency replacement - Broaden class declaration pattern to match final/abstract modifiers - Add 7 unit tests for add_codeflash_dependency_to_pom including stale system-scope replacement and reordered XML elements - Clarify comment about @SuppressWarnings in both modes Co-Authored-By: Claude Opus 4.6 --- codeflash/languages/java/build_tools.py | 14 ++- codeflash/languages/java/instrumentation.py | 8 +- .../test_java/test_build_tools.py | 104 ++++++++++++++++++ 3 files changed, 117 insertions(+), 9 deletions(-) diff --git a/codeflash/languages/java/build_tools.py b/codeflash/languages/java/build_tools.py index 571b6c2ba..61e484fa2 100644 --- a/codeflash/languages/java/build_tools.py +++ b/codeflash/languages/java/build_tools.py @@ -648,13 +648,15 @@ def add_codeflash_dependency_to_pom(pom_path: Path) -> bool: if "codeflash-runtime" in content: # If a previous run left a system-scope dependency, replace it with test scope. # System-scope dependencies cause Maven warnings and are rejected by some projects. - if "system" in content and "codeflash-runtime" in content: + if "system" in content: + # Match the codeflash-runtime dependency block regardless of XML element order + # within the tag, since Maven POMs don't enforce element ordering. content = re.sub( - r"\s*com\.codeflash\s*" - r"codeflash-runtime\s*" - r"[^<]*\s*" - r"system\s*" - r"[^<]*\s*", + r"\s*" + r"(?=(?:[\s\S]*?com\.codeflash))" + r"(?=(?:[\s\S]*?codeflash-runtime))" + r"(?=(?:[\s\S]*?system))" + r"[\s\S]*?", "\n" " com.codeflash\n" " codeflash-runtime\n" diff --git a/codeflash/languages/java/instrumentation.py b/codeflash/languages/java/instrumentation.py index d5af758c9..0f7b29610 100644 --- a/codeflash/languages/java/instrumentation.py +++ b/codeflash/languages/java/instrumentation.py @@ -640,8 +640,8 @@ def instrument_existing_test( # Add @SuppressWarnings("CheckReturnValue") to the class declaration. # Projects using Error Prone (e.g. Guava) enforce CheckReturnValue as a compiler error. - # Our instrumented tests intentionally discard return values in performance-only mode - # (after assertion stripping), which would fail compilation without this suppression. + # Applied in both modes: performance mode strips assertions (creating discarded return values), + # and behavior mode adds wrapper calls that may also discard return values. modified_source = _add_suppress_warnings_annotation(modified_source, new_class_name) # Add timing instrumentation to test methods @@ -841,7 +841,9 @@ def _add_suppress_warnings_annotation(source: str, class_name: str) -> str: Our instrumented tests intentionally discard return values after assertion stripping, which would fail compilation without this suppression. """ - class_decl_pattern = re.compile(rf"^((?:public\s+)?class\s+{re.escape(class_name)}\b)", re.MULTILINE) + class_decl_pattern = re.compile( + rf"^((?:(?:public|protected|final|abstract)\s+)*class\s+{re.escape(class_name)}\b)", re.MULTILINE + ) match = class_decl_pattern.search(source) if not match: return source diff --git a/tests/test_languages/test_java/test_build_tools.py b/tests/test_languages/test_java/test_build_tools.py index 5a194447e..f6429b16c 100644 --- a/tests/test_languages/test_java/test_build_tools.py +++ b/tests/test_languages/test_java/test_build_tools.py @@ -5,6 +5,7 @@ from codeflash.languages.java.build_tools import ( BuildTool, + add_codeflash_dependency_to_pom, detect_build_tool, find_maven_executable, find_source_root, @@ -454,3 +455,106 @@ def test_nonexistent_custom_dir_ignored(self, tmp_path): info = get_project_info(tmp_path) assert info is not None assert len(info.source_roots) == 1 + + +class TestAddCodeflashDependencyToPom: + """Tests for add_codeflash_dependency_to_pom, including stale system-scope replacement.""" + + def test_adds_dependency_to_clean_pom(self, tmp_path): + pom = tmp_path / "pom.xml" + pom.write_text( + '\n' + "\n" + " \n" + " \n" + " junit\n" + " junit\n" + " 4.13.2\n" + " \n" + " \n" + "\n", + encoding="utf-8", + ) + assert add_codeflash_dependency_to_pom(pom) is True + content = pom.read_text(encoding="utf-8") + assert "codeflash-runtime" in content + assert "test" in content + + def test_replaces_system_scope_with_test_scope(self, tmp_path): + pom = tmp_path / "pom.xml" + pom.write_text( + '\n' + "\n" + " \n" + " \n" + " com.codeflash\n" + " codeflash-runtime\n" + " 1.0.0\n" + " system\n" + " /some/path/jar.jar\n" + " \n" + " \n" + "\n", + encoding="utf-8", + ) + assert add_codeflash_dependency_to_pom(pom) is True + content = pom.read_text(encoding="utf-8") + assert "test" in content + assert "system" not in content + assert "" not in content + + def test_replaces_system_scope_with_reordered_elements(self, tmp_path): + """XML elements inside can appear in any order.""" + pom = tmp_path / "pom.xml" + pom.write_text( + '\n' + "\n" + " \n" + " \n" + " system\n" + " com.codeflash\n" + " /some/path/jar.jar\n" + " 1.0.0\n" + " codeflash-runtime\n" + " \n" + " \n" + "\n", + encoding="utf-8", + ) + assert add_codeflash_dependency_to_pom(pom) is True + content = pom.read_text(encoding="utf-8") + assert "test" in content + assert "system" not in content + assert "" not in content + + def test_skips_when_test_scope_already_present(self, tmp_path): + pom = tmp_path / "pom.xml" + pom.write_text( + '\n' + "\n" + " \n" + " \n" + " com.codeflash\n" + " codeflash-runtime\n" + " 1.0.0\n" + " test\n" + " \n" + " \n" + "\n", + encoding="utf-8", + ) + assert add_codeflash_dependency_to_pom(pom) is True + content = pom.read_text(encoding="utf-8") + assert content.count("codeflash-runtime") == 1 + + def test_returns_false_for_missing_pom(self, tmp_path): + pom = tmp_path / "pom.xml" + assert add_codeflash_dependency_to_pom(pom) is False + + def test_returns_false_when_no_dependencies_tag(self, tmp_path): + pom = tmp_path / "pom.xml" + pom.write_text( + '\n4.0.0\n', + encoding="utf-8", + ) + assert add_codeflash_dependency_to_pom(pom) is False From 62aaab87ac3f4d443a4eb2e4f56641e0909214a6 Mon Sep 17 00:00:00 2001 From: Mohamed Ashraf Date: Wed, 4 Mar 2026 00:37:32 +0000 Subject: [PATCH 3/4] fix: pre-install multi-module Maven deps to avoid recompilation failures Multi-module Maven projects like Guava fail on sequential Maven invocations because compiler plugin 3.15.0's JDK-8318913 workaround patches module-info.class timestamps, triggering unnecessary recompilation with -am that fails on partial reactor rebuilds. This pre-installs deps to .m2 once, then drops -am from all subsequent test commands. Co-Authored-By: Claude Opus 4.6 --- codeflash/languages/java/test_runner.py | 92 ++++++++++++++++-- tests/test_java_multimodule_deps_install.py | 102 ++++++++++++++++++++ 2 files changed, 188 insertions(+), 6 deletions(-) create mode 100644 tests/test_java_multimodule_deps_install.py diff --git a/codeflash/languages/java/test_runner.py b/codeflash/languages/java/test_runner.py index fd01d2623..c56a7d1bd 100644 --- a/codeflash/languages/java/test_runner.py +++ b/codeflash/languages/java/test_runner.py @@ -43,6 +43,10 @@ # so we avoid calling `mvn dependency:build-classpath` (~2-3s) repeatedly. _classpath_cache: dict[tuple[Path, str | None], str] = {} +# Cache for multi-module dependency installs — keyed on (maven_root, test_module). +# After pre-installing deps to .m2 once, subsequent Maven invocations can skip -am. +_multimodule_deps_installed: set[tuple[Path, str]] = set() + # Regex pattern for valid Java class names (package.ClassName format) # Allows: letters, digits, underscores, dots, and dollar signs (inner classes) _VALID_JAVA_CLASS_NAME = re.compile(r"^[a-zA-Z_$][a-zA-Z0-9_$.]*$") @@ -251,6 +255,68 @@ def _ensure_codeflash_runtime(maven_root: Path, test_module: str | None) -> bool return True +def ensure_multi_module_deps_installed(maven_root: Path, test_module: str | None, env: dict[str, str]) -> bool: + """Pre-install multi-module dependencies to the local Maven repository. + + In multi-module Maven projects (like Guava), Maven compiler plugin 3.15.0's + JDK-8318913 workaround patches module-info.class timestamps after compilation. + When a subsequent Maven invocation uses -am (also-make), the compiler detects + "changed source code" and recompiles dependency modules — which fails because + module-path resolution doesn't work in a partial reactor rebuild. + + This function runs `mvn install -DskipTests -pl -am` once to put all + dependency JARs into ~/.m2. After that, test-running commands can use + `-pl ` without `-am`, resolving deps from .m2 instead. + + Skipped for single-module projects (test_module is None) and cached so it only + runs once per (maven_root, test_module) pair within a session. + """ + if not test_module: + return True + + cache_key = (maven_root, test_module) + if cache_key in _multimodule_deps_installed: + logger.debug("Multi-module deps already installed for %s:%s, skipping", maven_root, test_module) + return True + + mvn = find_maven_executable() + if not mvn: + logger.error("Maven not found — cannot pre-install multi-module dependencies") + return False + + cmd = [ + mvn, + "install", + "-DskipTests", + "-B", + "-pl", + test_module, + "-am", + ] + cmd.extend(_MAVEN_VALIDATION_SKIP_FLAGS) + + logger.info("Pre-installing multi-module dependencies: %s (module: %s)", maven_root, test_module) + logger.debug("Running: %s", " ".join(cmd)) + + try: + result = _run_cmd_kill_pg_on_timeout(cmd, cwd=maven_root, env=env, timeout=300) + if result.returncode != 0: + logger.error( + "Failed to pre-install multi-module deps (exit %d).\nstdout: %s\nstderr: %s", + result.returncode, + result.stdout[-2000:] if result.stdout else "", + result.stderr[-2000:] if result.stderr else "", + ) + return False + except Exception: + logger.exception("Exception during multi-module dependency install") + return False + + _multimodule_deps_installed.add(cache_key) + logger.info("Multi-module dependencies installed successfully for %s:%s", maven_root, test_module) + return True + + def _extract_modules_from_pom_content(content: str) -> list[str]: """Extract module names from Maven POM XML content using proper XML parsing. @@ -485,6 +551,11 @@ def run_behavioral_tests( # Ensure codeflash-runtime is installed and added as dependency before compilation _ensure_codeflash_runtime(maven_root, test_module) + # Pre-install multi-module deps to .m2 so subsequent Maven runs don't need -am + base_env = os.environ.copy() + base_env.update(test_env) + ensure_multi_module_deps_installed(maven_root, test_module, base_env) + # Create SQLite database path for behavior capture - use standard path that parse_test_results expects sqlite_db_path = get_run_tmp_file(Path(f"test_return_values_{candidate_index}.sqlite")) @@ -604,7 +675,7 @@ def _compile_tests( cmd.extend(_MAVEN_VALIDATION_SKIP_FLAGS) if test_module: - cmd.extend(["-pl", test_module, "-am"]) + cmd.extend(["-pl", test_module]) logger.debug("Compiling tests: %s in %s", " ".join(cmd), project_root) @@ -1186,6 +1257,11 @@ def run_benchmarking_tests( # Ensure codeflash-runtime is installed and added as dependency before compilation _ensure_codeflash_runtime(maven_root, test_module) + # Pre-install multi-module deps to .m2 so subsequent Maven runs don't need -am + base_env = os.environ.copy() + base_env.update(test_env) + ensure_multi_module_deps_installed(maven_root, test_module, base_env) + # Get test class names test_classes = _get_test_class_names(test_paths, mode="performance") if not test_classes: @@ -1569,16 +1645,15 @@ def _run_maven_tests( if enable_coverage: cmd.append("-Dmaven.test.failure.ignore=true") - # For multi-module projects, specify which module to test + # For multi-module projects, specify which module to test. + # Dependencies are pre-installed to .m2 by ensure_multi_module_deps_installed(), + # so we use -pl without -am to avoid recompiling dependency modules (which fails + # on projects like Guava due to Maven compiler plugin's JDK-8318913 workaround). if test_module: - # -am = also make dependencies - # -DfailIfNoTests=false allows dependency modules without tests to pass - # -DskipTests=false overrides any skipTests=true in pom.xml cmd.extend( [ "-pl", test_module, - "-am", "-DfailIfNoTests=false", "-Dsurefire.failIfNoSpecifiedTests=false", "-DskipTests=false", @@ -2019,6 +2094,11 @@ def run_line_profile_tests( # Ensure codeflash-runtime is installed and added as dependency before compilation _ensure_codeflash_runtime(maven_root, test_module) + # Pre-install multi-module deps to .m2 so subsequent Maven runs don't need -am + base_env = os.environ.copy() + base_env.update(test_env) + ensure_multi_module_deps_installed(maven_root, test_module, base_env) + # Set up environment with profiling mode run_env = os.environ.copy() run_env.update(test_env) diff --git a/tests/test_java_multimodule_deps_install.py b/tests/test_java_multimodule_deps_install.py new file mode 100644 index 000000000..3a1390832 --- /dev/null +++ b/tests/test_java_multimodule_deps_install.py @@ -0,0 +1,102 @@ +"""Tests for ensure_multi_module_deps_installed in Java test runner.""" + +import subprocess +from pathlib import Path +from unittest.mock import MagicMock, patch + +import pytest + +from codeflash.languages.java.test_runner import ( + _multimodule_deps_installed, + ensure_multi_module_deps_installed, +) + + +@pytest.fixture(autouse=True) +def clear_cache(): + """Clear the multi-module deps cache before each test.""" + _multimodule_deps_installed.clear() + yield + _multimodule_deps_installed.clear() + + +def test_skipped_for_single_module(): + """Single-module projects (test_module=None) should be a no-op.""" + result = ensure_multi_module_deps_installed(Path("/fake"), None, {}) + assert result is True + assert len(_multimodule_deps_installed) == 0 + + +@patch("codeflash.languages.java.test_runner.find_maven_executable", return_value="mvn") +@patch("codeflash.languages.java.test_runner._run_cmd_kill_pg_on_timeout") +def test_runs_install_command_with_correct_args(mock_run, mock_mvn): + """Should run mvn install -DskipTests -pl -am with validation skip flags.""" + mock_run.return_value = subprocess.CompletedProcess(args=["mvn"], returncode=0, stdout="", stderr="") + + root = Path("/project") + result = ensure_multi_module_deps_installed(root, "guava-tests", {"JAVA_HOME": "/jdk"}) + + assert result is True + mock_run.assert_called_once() + cmd = mock_run.call_args[0][0] + assert cmd[0] == "mvn" + assert "install" in cmd + assert "-DskipTests" in cmd + assert "-pl" in cmd + assert "guava-tests" in cmd + assert "-am" in cmd + assert "-B" in cmd + # Validation skip flags should be present + assert "-Drat.skip=true" in cmd + assert "-Dcheckstyle.skip=true" in cmd + # cwd should be maven_root + assert mock_run.call_args[1]["cwd"] == root + + +@patch("codeflash.languages.java.test_runner.find_maven_executable", return_value="mvn") +@patch("codeflash.languages.java.test_runner._run_cmd_kill_pg_on_timeout") +def test_caches_and_does_not_rerun(mock_run, mock_mvn): + """Second call with same (root, module) should be cached — no Maven invocation.""" + mock_run.return_value = subprocess.CompletedProcess(args=["mvn"], returncode=0, stdout="", stderr="") + + root = Path("/project") + ensure_multi_module_deps_installed(root, "guava-tests", {}) + assert mock_run.call_count == 1 + + # Second call — should be cached + result = ensure_multi_module_deps_installed(root, "guava-tests", {}) + assert result is True + assert mock_run.call_count == 1 # NOT called again + + +@patch("codeflash.languages.java.test_runner.find_maven_executable", return_value="mvn") +@patch("codeflash.languages.java.test_runner._run_cmd_kill_pg_on_timeout") +def test_different_modules_not_cached(mock_run, mock_mvn): + """Different test modules should each trigger their own install.""" + mock_run.return_value = subprocess.CompletedProcess(args=["mvn"], returncode=0, stdout="", stderr="") + + root = Path("/project") + ensure_multi_module_deps_installed(root, "module-a", {}) + ensure_multi_module_deps_installed(root, "module-b", {}) + assert mock_run.call_count == 2 + + +@patch("codeflash.languages.java.test_runner.find_maven_executable", return_value="mvn") +@patch("codeflash.languages.java.test_runner._run_cmd_kill_pg_on_timeout") +def test_returns_false_on_maven_failure(mock_run, mock_mvn): + """Non-zero exit code should return False and NOT cache.""" + mock_run.return_value = subprocess.CompletedProcess( + args=["mvn"], returncode=1, stdout="", stderr="BUILD FAILURE" + ) + + root = Path("/project") + result = ensure_multi_module_deps_installed(root, "guava-tests", {}) + assert result is False + assert len(_multimodule_deps_installed) == 0 + + +@patch("codeflash.languages.java.test_runner.find_maven_executable", return_value=None) +def test_returns_false_when_maven_not_found(mock_mvn): + """Should return False if Maven executable is not found.""" + result = ensure_multi_module_deps_installed(Path("/fake"), "module", {}) + assert result is False From e4ceaff468499113c06b200e75ea5f5c589a373b Mon Sep 17 00:00:00 2001 From: Mohamed Ashraf Date: Wed, 4 Mar 2026 01:31:51 +0000 Subject: [PATCH 4/4] fix: prevent system-scope dependency regex from replacing all pom dependencies The regex for replacing system-scope codeflash-runtime dependencies used [\s\S]*? lookaheads that could scan forward across the entire file. This caused every block to match (since codeflash-runtime and system existed somewhere later in the file), replacing all 9 dependencies with codeflash-runtime duplicates. Fixed by using a callback function that matches each block individually and only replaces the one containing both "codeflash-runtime" and "system". Co-Authored-By: Claude Opus 4.6 --- codeflash/languages/java/build_tools.py | 35 ++++++++++++++----------- 1 file changed, 19 insertions(+), 16 deletions(-) diff --git a/codeflash/languages/java/build_tools.py b/codeflash/languages/java/build_tools.py index 61e484fa2..f10718415 100644 --- a/codeflash/languages/java/build_tools.py +++ b/codeflash/languages/java/build_tools.py @@ -649,22 +649,25 @@ def add_codeflash_dependency_to_pom(pom_path: Path) -> bool: # If a previous run left a system-scope dependency, replace it with test scope. # System-scope dependencies cause Maven warnings and are rejected by some projects. if "system" in content: - # Match the codeflash-runtime dependency block regardless of XML element order - # within the tag, since Maven POMs don't enforce element ordering. - content = re.sub( - r"\s*" - r"(?=(?:[\s\S]*?com\.codeflash))" - r"(?=(?:[\s\S]*?codeflash-runtime))" - r"(?=(?:[\s\S]*?system))" - r"[\s\S]*?", - "\n" - " com.codeflash\n" - " codeflash-runtime\n" - " 1.0.0\n" - " test\n" - " ", - content, - ) + # Replace ONLY the codeflash-runtime dependency block that has system scope. + # We find each ... block individually and only replace + # the one containing both "codeflash-runtime" and "system". + # The previous regex used [\s\S]*? lookaheads that could match across blocks, + # accidentally replacing every dependency in the file. + def replace_system_dep(match: re.Match) -> str: + block = match.group(0) + if "codeflash-runtime" in block and "system" in block: + return ( + "\n" + " com.codeflash\n" + " codeflash-runtime\n" + " 1.0.0\n" + " test\n" + " " + ) + return block + + content = re.sub(r"[\s\S]*?", replace_system_dep, content) pom_path.write_text(content, encoding="utf-8") logger.info("Replaced system-scope codeflash-runtime dependency with test scope") return True