fix: build patches from complete target context

This commit is contained in:
thanhnv
2026-07-18 13:44:20 +07:00
parent 850be145d7
commit c0a6dbc25f
2 changed files with 120 additions and 15 deletions
@@ -206,20 +206,34 @@ def context_candidates(project: dict, goal: str):
return sorted(candidates, key=lambda item: (-item[0], item[1]))
def build_context(job_path: str, project_id: str, goal: str):
def build_context(job_path: str, project_id: str, goal: str, write_intent=False):
project = registered_project(project_id)
candidates = context_candidates(project, goal)
excerpts, manifest_files, characters = [], [], 0
max_files, max_characters, per_file = 16, 7_000, 1_200
# A diff can only apply when the model sees the exact target-file content.
# Read-only analysis stays compact; write-intent gives the three most
# relevant files their complete bounded snapshot and retains supporting
# contracts within a larger, still finite context budget.
max_files = 12 if write_intent else 16
max_characters = 30_000 if write_intent else 7_000
for score, relative, raw in candidates:
if len(excerpts) >= max_files or characters >= max_characters:
break
excerpt = redact_context(raw[:min(per_file, max_characters - characters)]).strip()
per_file = 16_000 if write_intent and len(excerpts) < 3 else (3_000 if write_intent else 1_200)
source = redact_context(raw).strip()
excerpt = source[:min(per_file, max_characters - characters)].strip()
if not excerpt or context_excerpt_is_sensitive(excerpt):
continue
excerpts.append(f"### FILE: {relative}\n{excerpt}")
characters += len(excerpt)
manifest_files.append({"path": relative, "sha256": sha(raw), "characters": len(excerpt), "relevance": score})
manifest_files.append({
"path": relative,
"sha256": sha(raw),
"characters": len(excerpt),
"source_characters": len(source),
"truncated": len(excerpt) < len(source),
"relevance": score,
})
bundle = "\n\n".join(excerpts)
allowed, safe_bundle = scan(bundle, "input")
if not allowed:
@@ -238,7 +252,7 @@ def build_context(job_path: str, project_id: str, goal: str):
"files": manifest_files,
"file_count": len(manifest_files),
"characters": characters,
"truncated": len(manifest_files) < len(candidates),
"truncated": len(manifest_files) < len(candidates) or any(item["truncated"] for item in manifest_files),
"bundle_sha256": sha(bundle),
}
atomic_json(manifest_path, manifest)
@@ -284,6 +298,52 @@ def validate_patch_check(patch: str) -> None:
raise ValueError(f"goal_patch_check_failed:{detail}")
def normalize_unified_diff(patch: str) -> str:
"""Repair mechanical unified-diff hunk counts without changing code.
Language models frequently emit the intended +/- lines but miscalculate
the counts in `@@ -old,count +new,count @@`, or lose the single space on a
blank context line. Both make git report a syntactically corrupt patch.
CASAN can correct those two mechanical properties deterministically. It
never invents code lines, and leaves likely truncated hunks untouched.
"""
lines = patch.rstrip("\n").split("\n")
normalized, changed = [], False
hunk_pattern = re.compile(r"^@@ -(\d+)(?:,(\d+))? \+(\d+)(?:,(\d+))? @@(.*)$")
index = 0
while index < len(lines):
header = hunk_pattern.match(lines[index])
if not header:
normalized.append(lines[index])
index += 1
continue
body, cursor = [], index + 1
while cursor < len(lines) and not lines[cursor].startswith(("@@ ", "diff --git ")):
line = lines[cursor]
if line == "":
line = " "
changed = True
if not line.startswith((" ", "+", "-", "\\")):
return patch
body.append(line)
cursor += 1
old_count = sum(1 for line in body if line.startswith((" ", "-")))
new_count = sum(1 for line in body if line.startswith((" ", "+")))
declared_old = int(header.group(2) or "1")
declared_new = int(header.group(4) or "1")
# A replacement cut immediately after its '-' lines must not be
# reinterpreted as a valid deletion-only patch.
if cursor == len(lines) and body and body[-1].startswith("-") and declared_new > new_count:
return patch
corrected = f"@@ -{header.group(1)},{old_count} +{header.group(3)},{new_count} @@{header.group(5)}"
if corrected != lines[index] or declared_old != old_count or declared_new != new_count:
changed = True
normalized.append(corrected)
normalized.extend(body)
index = cursor
return "\n".join(normalized) + "\n" if changed else patch
def validate_write_output(text: str) -> str:
"""Fail at the producing harness when a write-intent reply is not a diff.
@@ -292,8 +352,17 @@ def validate_write_output(text: str) -> str:
that violated the required output contract.
"""
patch = extract_patch(text)
validate_patch_check(patch)
return patch
try:
validate_patch_check(patch)
return patch
except ValueError as original_error:
if "corrupt patch" not in str(original_error):
raise
normalized = normalize_unified_diff(patch)
if normalized == patch:
raise
validate_patch_check(normalized)
return normalized
def patch_repair_attempts() -> int:
@@ -381,7 +450,7 @@ def repair_write_output(model: str, original_prompt: str, invalid_output: str, c
})
continue
try:
validate_write_output(output)
checked_output = validate_write_output(output)
except ValueError as error:
last_reason = str(error)
attempts.append({
@@ -408,7 +477,7 @@ def repair_write_output(model: str, original_prompt: str, invalid_output: str, c
"reason": "ok",
})
metadata["repair_attempts"] = attempts
return True, output, metadata, "ok"
return True, checked_output, metadata, "ok"
last_metadata["repair_attempts"] = attempts
return False, "", last_metadata, f"goal_patch_repair_invalid:{last_reason}"
@@ -762,7 +831,8 @@ def run(job_path: str) -> int:
if not allowed:
emit(goal_id, "H4-security", "blocked", "Objective rejected by security boundary")
raise ValueError("goal_security_blocked")
context_bundle, context_manifest, context_manifest_path = build_context(job_path, project_id, safe_goal)
write_intent = requests_side_effect(safe_goal)
context_bundle, context_manifest, context_manifest_path = build_context(job_path, project_id, safe_goal, write_intent)
context_summary = {
"files": context_manifest["file_count"],
"characters": context_manifest["characters"],
@@ -776,8 +846,6 @@ def run(job_path: str) -> int:
})
emit(goal_id, "H4-security", "running", "Objective and workspace snapshot passed; model outputs pending")
write_intent = requests_side_effect(safe_goal)
stage(job_path, "local-worker", "running", "Local model is developing the primary solution", job.get("local_provider", ""), local_model)
emit(goal_id, "H2-tool", "running", "Local worker is developing a solution", {"provider": job.get("local_provider", ""), "model": local_model})
output_contract = (
@@ -802,7 +870,7 @@ def run(job_path: str) -> int:
raise ValueError("local_output_security_blocked")
if write_intent:
try:
validate_write_output(safe_local)
safe_local = validate_write_output(safe_local)
except ValueError as first_error:
repair_candidates = patch_repair_models(local_model)[:patch_repair_attempts()]
repair_model = repair_candidates[0] if repair_candidates else local_model
@@ -819,7 +887,7 @@ def run(job_path: str) -> int:
emit(goal_id, "H4-security", "blocked", "Repaired local worker output rejected")
raise ValueError("local_output_security_blocked")
try:
validate_write_output(safe_local)
safe_local = validate_write_output(safe_local)
reason = ""
except ValueError as repair_error:
reason = f"goal_patch_repair_invalid:{repair_error}"
@@ -866,7 +934,7 @@ def run(job_path: str) -> int:
raise ValueError("cloud_output_security_blocked")
if write_intent:
try:
validate_write_output(safe_result)
safe_result = validate_write_output(safe_result)
except ValueError as error:
cloud_ok = False
cloud_reason = f"reviewer_output_contract_invalid:{error}"
@@ -43,6 +43,43 @@ class GoalPatchWorkflowTests(unittest.TestCase):
with self.assertRaisesRegex(ValueError, "goal_patch_check_failed"):
ORCHESTRATOR.validate_write_output(truncated)
def test_hunk_counts_and_blank_context_are_repaired_deterministically(self):
corrupt = (
"diff --git a/apps/okr/frontend/a.ts b/apps/okr/frontend/a.ts\n"
"--- a/apps/okr/frontend/a.ts\n+++ b/apps/okr/frontend/a.ts\n"
"@@ -10,8 +10,9 @@\n context\n-old\n+new\n\n tail\n"
)
normalized = ORCHESTRATOR.normalize_unified_diff(corrupt)
self.assertIn("@@ -10,4 +10,4 @@", normalized)
self.assertIn("\n \n tail\n", normalized)
def test_write_validation_uses_normalized_patch_after_corrupt_error(self):
corrupt = "diff --git a/a b/a\n--- a/a\n+++ b/a\n@@ -1,9 +1,9 @@\n-old\n+new\n"
with patch.object(ORCHESTRATOR, "validate_patch_check", side_effect=[ValueError("goal_patch_check_failed:error: corrupt patch at line 7"), None]) as check:
normalized = ORCHESTRATOR.validate_write_output(corrupt)
self.assertIn("@@ -1,1 +1,1 @@", normalized)
self.assertEqual(check.call_count, 2)
def test_truncated_replacement_is_not_reinterpreted_as_deletion(self):
truncated = "diff --git a/a b/a\n--- a/a\n+++ b/a\n@@ -1,1 +1,1 @@\n-old\n"
self.assertEqual(ORCHESTRATOR.normalize_unified_diff(truncated), truncated)
def test_write_context_keeps_the_most_relevant_target_file_complete(self):
raw = "export function KeyResultDetail() {\n" + (" return null;\n" * 180) + "}\n"
candidates = [
(20, "apps/okr/frontend/src/pages/KeyResultDetail.tsx", raw),
(5, "docs/technical_architecture.md", "architecture" * 400),
]
project = {"domain": "OKR", "domain_root": "apps/okr", "roots": []}
with tempfile.TemporaryDirectory() as directory, \
patch.object(ORCHESTRATOR, "registered_project", return_value=project), \
patch.object(ORCHESTRATOR, "context_candidates", return_value=candidates), \
patch.object(ORCHESTRATOR, "scan", side_effect=lambda text, mode: (True, text)):
bundle, manifest, _ = ORCHESTRATOR.build_context(os.path.join(directory, "goal.json"), "okr", "KeyResultDetail", True)
self.assertIn(raw.strip(), bundle)
self.assertEqual(manifest["files"][0]["characters"], len(raw.strip()))
self.assertFalse(manifest["files"][0]["truncated"])
def test_invalid_write_output_gets_one_bounded_repair_attempt(self):
repaired = "```diff\ndiff --git a/apps/okr/frontend/a.ts b/apps/okr/frontend/a.ts\n--- a/apps/okr/frontend/a.ts\n+++ b/apps/okr/frontend/a.ts\n@@ -1 +1 @@\n-a\n+b\n```"
with patch.dict(os.environ, {"CASAN_GOAL_PATCH_REPAIR_ATTEMPTS": "1"}, clear=False), \