Skip to content

Commit 10ebf38

Browse files
committed
fix: address PR review findings in testgen review/repair loop
- Revert repaired test files on disk before returning Failure when post-repair re-validation fails entirely - Break out of repair cycle loop when all repair API calls fail instead of wasting cycles retrying - Fix --testgen-review-turns help text to say default: 2 (matches MAX_TEST_REPAIR_CYCLES)
1 parent f35898d commit 10ebf38

2 files changed

Lines changed: 72 additions & 64 deletions

File tree

codeflash/cli_cmds/cli.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -110,7 +110,7 @@ def parse_args() -> Namespace:
110110
"--testgen-review", default=False, action="store_true", help="Enable AI review and repair of generated tests"
111111
)
112112
parser.add_argument(
113-
"--testgen-review-turns", type=int, default=None, help="Number of review/repair cycles (default: 1)"
113+
"--testgen-review-turns", type=int, default=None, help="Number of review/repair cycles (default: 2)"
114114
)
115115
parser.add_argument(
116116
"--async",

codeflash/optimization/function_optimizer.py

Lines changed: 71 additions & 63 deletions
Original file line numberDiff line numberDiff line change
@@ -2138,74 +2138,82 @@ def review_and_repair_tests(
21382138
repaired_files += 1
21392139
repaired_indices.add(review.test_index)
21402140

2141-
if any_repaired:
2142-
generated_tests = self.language_support.postprocess_generated_tests(
2143-
generated_tests,
2144-
test_framework=self.test_cfg.test_framework,
2145-
project_root=self.project_root,
2146-
source_file_path=self.function_to_optimize.file_path,
2147-
)
2148-
console.print(f" [green]Repaired {repaired_files} test file(s)[/green]")
2149-
with progress_bar("Re-validating repaired tests..."):
2150-
validation = self.run_behavioral_validation(
2151-
code_context, original_helper_code, file_path_to_helper_classes
2152-
)
2153-
if validation is None:
2154-
return Failure("Repaired tests failed behavioral validation.")
2155-
behavioral_results, coverage_results = validation
2156-
2157-
# Check which repaired test files still have failures and revert them
2158-
still_failing_files: set[Path] = set()
2159-
for result in behavioral_results.test_results:
2160-
if result.test_type == TestType.GENERATED_REGRESSION and not result.did_pass:
2161-
still_failing_files.add(result.file_name)
2141+
if not any_repaired:
2142+
break
21622143

2163-
reverted_indices = set()
2144+
generated_tests = self.language_support.postprocess_generated_tests(
2145+
generated_tests,
2146+
test_framework=self.test_cfg.test_framework,
2147+
project_root=self.project_root,
2148+
source_file_path=self.function_to_optimize.file_path,
2149+
)
2150+
console.print(f" [green]Repaired {repaired_files} test file(s)[/green]")
2151+
with progress_bar("Re-validating repaired tests..."):
2152+
validation = self.run_behavioral_validation(
2153+
code_context, original_helper_code, file_path_to_helper_classes
2154+
)
2155+
if validation is None:
21642156
for idx in repaired_indices:
21652157
gt = generated_tests.generated_tests[idx]
2166-
if gt.behavior_file_path in still_failing_files:
2167-
orig_source, orig_behavior, orig_perf, orig_raw = pre_repair_snapshots[idx]
2168-
gt.generated_original_test_source = orig_source
2169-
gt.instrumented_behavior_test_source = orig_behavior
2170-
gt.instrumented_perf_test_source = orig_perf
2171-
gt.raw_generated_test_source = orig_raw
2172-
gt.behavior_file_path.write_text(orig_behavior, encoding="utf8")
2173-
gt.perf_file_path.write_text(orig_perf, encoding="utf8")
2174-
reverted_indices.add(idx)
2175-
2176-
# Show diffs only for repairs that survived re-validation
2177-
successful_repairs = [r for r in all_to_repair if r.test_index not in reverted_indices]
2178-
if successful_repairs:
2179-
self.display_repaired_functions(generated_tests, successful_repairs, original_sources)
2180-
2181-
if reverted_indices:
2182-
console.print(
2183-
f" [yellow]Reverted {len(reverted_indices)} test file(s) "
2184-
f"that still failed after repair[/yellow]"
2185-
)
2186-
# Collect error messages from failed repairs so the next cycle can learn from them
2187-
revalidation_failures = behavioral_results.test_failures or {}
2188-
for idx in reverted_indices:
2189-
gt = generated_tests.generated_tests[idx]
2190-
errors_for_file: dict[str, str] = {}
2191-
for result in behavioral_results.test_results:
2192-
if (
2193-
result.file_name == gt.behavior_file_path
2194-
and result.test_type == TestType.GENERATED_REGRESSION
2195-
and not result.did_pass
2196-
and result.id.test_function_name
2197-
):
2198-
fn_name = result.id.test_fn_qualified_name()
2199-
errors_for_file[fn_name] = revalidation_failures.get(fn_name, "Test failed")
2200-
if errors_for_file:
2201-
previous_repair_errors[idx] = errors_for_file
2202-
# Invalidate behavioral results since we reverted some files
2203-
behavioral_results = None
2204-
coverage_results = None
2158+
orig_source, orig_behavior, orig_perf, orig_raw = pre_repair_snapshots[idx]
2159+
gt.generated_original_test_source = orig_source
2160+
gt.instrumented_behavior_test_source = orig_behavior
2161+
gt.instrumented_perf_test_source = orig_perf
2162+
gt.raw_generated_test_source = orig_raw
2163+
gt.behavior_file_path.write_text(orig_behavior, encoding="utf8")
2164+
gt.perf_file_path.write_text(orig_perf, encoding="utf8")
2165+
return Failure("Repaired tests failed behavioral validation.")
2166+
behavioral_results, coverage_results = validation
2167+
2168+
# Check which repaired test files still have failures and revert them
2169+
still_failing_files: set[Path] = set()
2170+
for result in behavioral_results.test_results:
2171+
if result.test_type == TestType.GENERATED_REGRESSION and not result.did_pass:
2172+
still_failing_files.add(result.file_name)
2173+
2174+
reverted_indices = set()
2175+
for idx in repaired_indices:
2176+
gt = generated_tests.generated_tests[idx]
2177+
if gt.behavior_file_path in still_failing_files:
2178+
orig_source, orig_behavior, orig_perf, orig_raw = pre_repair_snapshots[idx]
2179+
gt.generated_original_test_source = orig_source
2180+
gt.instrumented_behavior_test_source = orig_behavior
2181+
gt.instrumented_perf_test_source = orig_perf
2182+
gt.raw_generated_test_source = orig_raw
2183+
gt.behavior_file_path.write_text(orig_behavior, encoding="utf8")
2184+
gt.perf_file_path.write_text(orig_perf, encoding="utf8")
2185+
reverted_indices.add(idx)
2186+
2187+
# Show diffs only for repairs that survived re-validation
2188+
successful_repairs = [r for r in all_to_repair if r.test_index not in reverted_indices]
2189+
if successful_repairs:
2190+
self.display_repaired_functions(generated_tests, successful_repairs, original_sources)
2191+
2192+
if reverted_indices:
2193+
console.print(
2194+
f" [yellow]Reverted {len(reverted_indices)} test file(s) that still failed after repair[/yellow]"
2195+
)
2196+
# Collect error messages from failed repairs so the next cycle can learn from them
2197+
revalidation_failures = behavioral_results.test_failures or {}
2198+
for idx in reverted_indices:
2199+
gt = generated_tests.generated_tests[idx]
2200+
errors_for_file: dict[str, str] = {}
2201+
for result in behavioral_results.test_results:
2202+
if (
2203+
result.file_name == gt.behavior_file_path
2204+
and result.test_type == TestType.GENERATED_REGRESSION
2205+
and not result.did_pass
2206+
and result.id.test_function_name
2207+
):
2208+
fn_name = result.id.test_fn_qualified_name()
2209+
errors_for_file[fn_name] = revalidation_failures.get(fn_name, "Test failed")
2210+
if errors_for_file:
2211+
previous_repair_errors[idx] = errors_for_file
2212+
# Invalidate behavioral results since we reverted some files
2213+
behavioral_results = None
2214+
coverage_results = None
22052215

22062216
console.rule()
2207-
# When all repair API calls failed (any_repaired=False), behavioral_results are from before
2208-
# the repair attempts. This is correct since no test code actually changed.
22092217
return Success((generated_tests, behavioral_results, coverage_results))
22102218

22112219
def find_and_process_best_optimization(

0 commit comments

Comments
 (0)