-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathFixGenerator.py
More file actions
475 lines (410 loc) · 19.1 KB
/
Copy pathFixGenerator.py
File metadata and controls
475 lines (410 loc) · 19.1 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
"""
FixGenerator.py — Generate an LLM bug fix as a unified diff.
Dataset- and infrastructure-agnostic: it knows nothing about Defects4J or Docker.
The caller (e.g. ``Experiment.py``) is responsible for locating and reading the
buggy sources, applying the resulting diff and validating it. FixGenerator only:
1. Builds a prompt from the bug description and the buggy source contents it is
handed.
2. Queries the LLM and extracts the diff from the response.
3. Optionally persists the generation artifacts (``fix.diff`` and the raw
response) and returns the generation metadata.
"""
import difflib
import os
import re
import time
from datetime import datetime, timezone
from llms import GoogleLLM, OpenAILLM, OpenRouterLLM, OllamaLLM, CopilotLLM, AnthropicLLM
# --- SEARCH/REPLACE block handling -------------------------------------------
#
# Rather than ask the model for a unified diff directly — which forces it to
# reproduce exact line numbers and context lines, something weaker models
# routinely hallucinate — we ask for Aider-style SEARCH/REPLACE blocks and build
# the diff ourselves with ``difflib`` against the real source. This removes every
# context-anchoring failure mode: the model only has to name the code to change
# and what to change it to; we do the anchoring against ground-truth text.
_SEARCH_REPLACE_RE = re.compile(
r"<{5,9} *SEARCH[^\n]*\n(?P<search>.*?)\n?={5,9}[^\n]*\n(?P<replace>.*?)\n?>{5,9} *REPLACE",
re.DOTALL,
)
def parse_search_replace_blocks(response_text):
"""Parse Aider-style SEARCH/REPLACE blocks from an LLM response.
Returns a list of ``(path, search, replace)`` tuples. ``path`` is taken from
the last non-empty, non-fence line preceding each block (empty if none).
"""
blocks = []
for match in _SEARCH_REPLACE_RE.finditer(response_text):
prefix_lines = response_text[: match.start()].split("\n")
path = ""
for line in reversed(prefix_lines):
stripped = line.strip()
if not stripped or stripped.startswith("```"):
continue
# Tolerate a leading "File:" / "path:" label.
path = re.sub(r"^(?:file|path)\s*:\s*", "", stripped, flags=re.IGNORECASE)
path = path.strip().strip("`").strip()
break
blocks.append((path, match.group("search"), match.group("replace")))
return blocks
def _normalize_for_match(line):
"""Collapse all whitespace so matching tolerates reformatting (e.g.
``for (`` vs ``for(``, changed indentation)."""
return "".join(line.split())
def _resolve_source_path(block_path, source_paths):
"""Map a block's declared path to one of the real source paths."""
if block_path in source_paths:
return block_path
base = os.path.basename(block_path)
for rel in source_paths:
if block_path and (rel.endswith(block_path) or os.path.basename(rel) == base):
return rel
if len(source_paths) == 1:
return source_paths[0]
return None
def _apply_block(content, search, replace):
"""Replace the first whitespace-insensitive match of ``search`` in
``content``. Returns ``(new_content, applied)``."""
file_lines = content.split("\n")
search_lines = search.split("\n")
replace_lines = replace.split("\n")
if not search or not search_lines:
return content, False
norm_file = [_normalize_for_match(l) for l in file_lines]
norm_search = [_normalize_for_match(l) for l in search_lines]
span = len(norm_search)
for i in range(len(file_lines) - span + 1):
if norm_file[i:i + span] == norm_search:
new_lines = file_lines[:i] + replace_lines + file_lines[i + span:]
return "\n".join(new_lines), True
return content, False
def _unified_file_diff(rel_path, original, patched):
"""Produce a git-appliable unified diff between two file contents."""
diff = difflib.unified_diff(
original.split("\n"),
patched.split("\n"),
fromfile=f"a/{rel_path}",
tofile=f"b/{rel_path}",
lineterm="",
)
return "\n".join(diff)
def build_diff_from_blocks(sources, blocks):
"""Apply SEARCH/REPLACE ``blocks`` to ``sources`` and return a unified diff.
Args:
sources: List of ``(relative_path, content)``.
blocks: List of ``(path, search, replace)`` from
:func:`parse_search_replace_blocks`.
Returns:
``(diff_text, applied, failed)`` where ``diff_text`` is the combined
unified diff for every changed file, ``applied`` is the number of blocks
that matched, and ``failed`` is a list of ``(path, reason)`` for blocks
that could not be located.
"""
originals = {rel: content for rel, content in sources}
working = dict(originals)
source_paths = list(originals)
applied = 0
failed = []
for block_path, search, replace in blocks:
rel = _resolve_source_path(block_path, source_paths)
if rel is None:
failed.append((block_path, "no matching source file"))
continue
new_content, ok = _apply_block(working[rel], search, replace)
if ok:
working[rel] = new_content
applied += 1
else:
failed.append((block_path or rel, "SEARCH text not found in source"))
diffs = [
_unified_file_diff(rel, originals[rel], working[rel])
for rel in source_paths
if working[rel] != originals[rel]
]
return "\n".join(d for d in diffs if d), applied, failed
# Matches a unified-diff hunk header, capturing the old/new start lines and any
# trailing section heading (e.g. the enclosing function name git echoes). The
# line counts are intentionally not captured: ``normalize_diff`` recomputes them.
_HUNK_HEADER_RE = re.compile(r"^@@ -(\d+)(?:,\d+)? \+(\d+)(?:,\d+)? @@(.*)$")
_FILE_HEADER_PREFIXES = ("--- ", "+++ ", "diff ", "index ")
def normalize_diff(diff_text: str) -> str:
"""Repair common malformations in LLM-generated unified diffs.
Three fixes are applied, all of which frequently break ``git apply``/
``patch`` on otherwise-correct LLM output:
1. Blank context lines emitted without their leading space are rewritten as
single-space context lines (otherwise the hunk is treated as ending
prematurely: "corrupt patch" / "unexpectedly ends in middle of line").
2. A trailing marker or comment appended after the last hunk (e.g.
``*** End of File ***``), despite explicit instructions not to, is
dropped: any in-hunk line that isn't a context (" "), removal ("-"),
addition ("+") or no-newline-marker ("\\") line ends the diff.
3. Each ``@@`` hunk header's line counts are recomputed from the actual
hunk body. Models routinely emit wrong counts, which makes the parser
mis-detect the hunk boundary and reject the whole patch as corrupt.
"""
cleaned = _strip_hunk_bodies(diff_text.split("\n"))
return "\n".join(_recount_hunk_headers(cleaned))
def _strip_hunk_bodies(lines):
"""Fix blank context lines and drop trailing non-diff garbage."""
out = []
in_hunk = False
for line in lines:
if line.startswith("@@"):
in_hunk = True
out.append(line)
continue
# A new file header ends the current hunk body.
if line.startswith(_FILE_HEADER_PREFIXES):
in_hunk = False
out.append(line)
continue
if in_hunk and line == "":
# Blank context line that lost its leading space.
out.append(" ")
continue
if in_hunk and line[:1] not in (" ", "+", "-", "\\"):
# Trailing garbage after the last hunk line; the diff is over.
break
out.append(line)
return out
def _recount_hunk_headers(lines):
"""Rewrite every ``@@`` header's line counts to match its body.
Headers without parseable start lines (e.g. a bare ``@@ @@``) are left
untouched, since their offsets can't be recovered here.
"""
out = []
i = 0
n = len(lines)
while i < n:
line = lines[i]
match = _HUNK_HEADER_RE.match(line)
if not match:
out.append(line)
i += 1
continue
old_start, new_start, heading = match.groups()
body = []
j = i + 1
while j < n:
body_line = lines[j]
if body_line.startswith("@@") or body_line.startswith(_FILE_HEADER_PREFIXES):
break
body.append(body_line)
j += 1
old_count = sum(1 for b in body if b[:1] in (" ", "-"))
new_count = sum(1 for b in body if b[:1] in (" ", "+"))
out.append(f"@@ -{old_start},{old_count} +{new_start},{new_count} @@{heading}")
out.extend(body)
i = j
return out
class FixGenerator:
"""Generates an LLM fix (unified diff) from a bug description and sources."""
# Generation budget. Both defaults were raised after the first campaign, in
# which they silently decided 46 runs' outcomes: 30 stopped at the old
# 24576-token output cap and 16 exhausted a 49152-token context window,
# every one of them recorded as the model failing to fix the bug.
#
# Sized from the campaign itself, not guessed: the largest successful
# generation used 18595 output tokens (reasoning included -- Ollama's
# eval_count counts both), so 32768 leaves ~75% headroom over the largest
# answer ever produced.
#
# 131072 is not a round number chosen for headroom: it is gpt-oss:120b's
# *native* context length (qwen3.6:35b's is 262144). Both models must get
# the same window or the comparison reacquires the asymmetry this exists to
# remove, and asking gpt-oss for more would be clamped silently -- the exact
# class of defect being fixed. So the shared window is the smaller ceiling.
# Measured to fit with room to spare: qwen 24830 MiB on a 46 GB L40S,
# gpt-oss 61700 MiB on a 96 GB H100 (scripts/probeContextVram.sh).
#
# RAISING THIS ABOVE 131072 IS NOT SAFE without re-checking both models'
# native context, and scripts/ollama_serve.sh must stay in step -- it reads
# this constant rather than repeating it.
DEFAULT_MAX_TOKENS = 32768
DEFAULT_CONTEXT_LENGTH = 131072
def __init__(self, model="ollama/gpt-oss:20b", temperature=0.0,
max_tokens=DEFAULT_MAX_TOKENS,
context_length=DEFAULT_CONTEXT_LENGTH):
self.model = model
self.temperature = temperature
self.max_tokens = max_tokens
self.context_length = context_length
self.llm = self._initialize_llm()
# ------------------------------------------------------------------ LLM
def _initialize_llm(self):
"""Initialize the LLM client based on the model identifier."""
providers = [OllamaLLM, CopilotLLM, AnthropicLLM, OpenRouterLLM, GoogleLLM, OpenAILLM]
for provider in providers:
if provider.is_supported(self.model):
# Ollama defaults to free-form text output (response_format=None),
# which is what we need for a unified diff.
#
# `context_length` is Ollama-specific (the hosted providers size
# their own window), so it is offered as a keyword and the
# providers that do not take it are called as before rather than
# all six having to grow a parameter they ignore.
try:
return provider.initialize(
self.model, self.temperature, self.max_tokens,
context_length=self.context_length,
)
except TypeError:
return provider.initialize(
self.model, self.temperature, self.max_tokens
)
raise ValueError(f"No provider found for model: {self.model}")
# -------------------------------------------------------------- prompt
def _build_prompt(self, sources, test_sources=None, test_log=None, issue_text=None):
"""Build the fix-generation prompt.
Args:
sources: List of ``(relative_path, content)`` for the buggy file(s).
test_sources: Optional list of ``(relative_path, content)`` for the
regression (trigger) test file(s).
test_log: Optional free-form text with the regression test's
failure output.
issue_text: Optional free-form text with the original bug-tracker
issue report.
"""
source_sections = []
for rel_path, content in sources:
source_sections.append(
f"--- FILE: {rel_path} ---\n{content}"
)
sources_block = "\n\n".join(source_sections)
optional_sections = ""
if test_sources:
test_source_sections = []
for rel_path, content in test_sources:
test_source_sections.append(
f"--- FILE: {rel_path} ---\n{content}"
)
optional_sections += (
"\n[REGRESSION TEST SOURCE FILE(S)]\n"
+ "\n\n".join(test_source_sections)
+ "\n"
)
if test_log:
optional_sections += f"\n[REGRESSION TEST LOG]\n{test_log}\n"
if issue_text:
optional_sections += f"\n[ORIGINAL ISSUE REPORT]\n{issue_text}\n"
first_path = sources[0][0] if sources else "path/To/File.java"
return f"""[SYSTEM INSTRUCTION]
You are an expert software engineer fixing a real bug. You will be given the \
buggy source file(s). Produce a correct, complete fix.
[BUGGY SOURCE FILE(S)]
{sources_block}
{optional_sections}
[TASK]
Fix the bug so that the failing (triggering) tests pass while keeping all other \
tests passing. Make all changes necessary to fix the bug correctly. Do not alter \
test files.
[OUTPUT FORMAT]
Do NOT output a diff. Describe every change as one or more *SEARCH/REPLACE blocks*.
For each change, write the path of the file to edit on its own line — exactly as \
shown above (e.g. `{first_path}`) — followed immediately by a block in this shape:
<<<<<<< SEARCH
(lines copied VERBATIM from the file shown above)
=======
(the replacement lines)
>>>>>>> REPLACE
Rules:
- The SEARCH section MUST be an exact, contiguous copy of lines from the file \
shown above: same characters, same indentation. Do NOT paraphrase, reformat, \
renumber, add, or omit lines. If it is not a verbatim copy it cannot be located \
and the change is rejected.
- Include enough surrounding lines (aim for 3 or more) so the SEARCH section \
matches exactly one place in the file.
- Keep each block small and focused; use several blocks instead of one large one.
- To insert code, copy an existing anchor into SEARCH and repeat it plus the new \
lines in REPLACE.
- Output ONLY file paths and SEARCH/REPLACE blocks — no diff, no line numbers, no \
explanations, no markdown code fences.
[EXAMPLE]
{first_path}
<<<<<<< SEARCH
if (hexDigits > 8) {{ // too many for an int
return createLong(str);
}}
return createInteger(str);
=======
if (hexDigits > 8) {{ // too many for an int
return createLong(str);
}}
if (hexDigits == 8 && firstDigit >= '8') {{
return createLong(str);
}}
return createInteger(str);
>>>>>>> REPLACE
"""
# ----------------------------------------------------------------- run
def generate(self, sources, test_sources=None, test_log=None,
issue_text=None, results_dir=None):
"""Generate a fix from the buggy sources.
Args:
sources: List of ``(relative_path, content)`` for the buggy file(s).
test_sources: Optional list of ``(relative_path, content)`` for the
regression (trigger) test file(s).
test_log: Optional free-form text with the regression test's
failure output.
issue_text: Optional free-form text with the original bug-tracker
issue report.
results_dir: Optional directory; when given, the generation artifacts
(``prompt.txt``, ``fix.diff`` and the raw response) are written
there.
Returns:
A dict with the generation metadata: ``model``, ``temperature``,
``timestamp``, ``elapsed_seconds``, ``prompt``, ``diff``,
``raw_response``, ``usage_metadata`` and the SEARCH/REPLACE
bookkeeping (``blocks_parsed``, ``blocks_applied``, ``blocks_failed``).
Validation (applying the diff, running tests) is the caller's
responsibility.
"""
timestamp = datetime.now(timezone.utc).isoformat()
prompt = self._build_prompt(sources, test_sources, test_log, issue_text)
print(f"[fixgen] Querying LLM ({self.model}) for a fix ...")
start = time.time()
response = self.llm.invoke(prompt)
elapsed = round(time.time() - start, 3)
# The model returns SEARCH/REPLACE blocks; we anchor them against the
# real source and build the unified diff ourselves with difflib.
blocks = parse_search_replace_blocks(response.content)
diff_text, applied, failed = build_diff_from_blocks(sources, blocks)
print(
f"[fixgen] SEARCH/REPLACE blocks: {len(blocks)} parsed, "
f"{applied} applied, {len(failed)} failed."
)
for path, reason in failed:
print(f"[fixgen] WARNING: block for {path!r} skipped: {reason}")
if results_dir is not None:
self._write_text(results_dir, "prompt.txt", prompt)
self._write_text(results_dir, "fix.diff", diff_text)
self._write_text(results_dir, "raw_response.txt", response.content)
return {
"model": self.model,
"temperature": self.temperature,
"timestamp": timestamp,
"elapsed_seconds": elapsed,
"prompt": prompt,
"diff": diff_text,
"raw_response": response.content,
"usage_metadata": getattr(response, "usage_metadata", None),
"blocks_parsed": len(blocks),
"blocks_applied": applied,
"blocks_failed": failed,
# How the generation ended. Without this an empty patch is
# ambiguous between "the model had no fix", "it was cut off at the
# token ceiling" and "its answer was never read" -- three very
# different facts that a repair benchmark must not average together.
"generation_status": getattr(response, "generation_status",
lambda: "ok")(),
"done_reason": getattr(response, "done_reason", ""),
"response_truncated": bool(getattr(response, "truncated", False)),
"context_exhausted": bool(getattr(response, "context_exhausted",
False)),
"reasoning_chars": getattr(response, "thinking_chars", 0),
}
@staticmethod
def _write_text(results_dir, filename, content):
os.makedirs(results_dir, exist_ok=True)
path = os.path.join(results_dir, filename)
with open(path, "w", encoding="utf-8") as f:
f.write(content or "")