fix: prioritize direct providers for goal patch recovery
This commit is contained in:
@@ -286,6 +286,29 @@ def patch_repair_attempts() -> int:
|
||||
return min(max(configured, 0), 1)
|
||||
|
||||
|
||||
def patch_repair_models(primary_model: str):
|
||||
"""Use direct cloud -> gateway -> local order for one bounded H2 recovery."""
|
||||
configured = [item.strip() for item in os.environ.get("CASAN_GOAL_PATCH_REPAIR_MODELS", "").split(",") if item.strip()]
|
||||
unique, seen = [], set()
|
||||
for candidate in configured + [primary_model]:
|
||||
if candidate and candidate not in seen:
|
||||
seen.add(candidate)
|
||||
unique.append(candidate)
|
||||
return unique
|
||||
|
||||
|
||||
def provider_for_model(model: str) -> str:
|
||||
if model.startswith("openai:"):
|
||||
return "openai"
|
||||
if model.startswith("anthropic:"):
|
||||
return "anthropic"
|
||||
if model.startswith("openai-compatible:"):
|
||||
return "omniroute"
|
||||
if model.startswith("ollama:"):
|
||||
return "ollama"
|
||||
return "model"
|
||||
|
||||
|
||||
def repair_write_output(model: str, original_prompt: str, invalid_output: str):
|
||||
"""Ask the producing model to repair format only, without relaxing H2.
|
||||
|
||||
@@ -303,7 +326,11 @@ def repair_write_output(model: str, original_prompt: str, invalid_output: str):
|
||||
"Do not explain, plan, summarize, use placeholders, or perform side effects. Preserve the original objective and workspace restrictions.\n\n"
|
||||
f"ORIGINAL CONTRACT:\n{original_prompt}\n\nPREVIOUS INVALID RESPONSE:\n{invalid_output[:8000]}"
|
||||
)
|
||||
ok, output, metadata, reason = call_model(model, repair_prompt, False)
|
||||
candidates = patch_repair_models(model)
|
||||
candidate = candidates[0] if candidates else model
|
||||
ok, output, metadata, reason = call_model(candidate, repair_prompt, candidate.startswith(("openai:", "anthropic:", "openai-compatible:")))
|
||||
metadata = dict(metadata)
|
||||
metadata["repair_model"] = candidate
|
||||
if not ok:
|
||||
return False, "", metadata, f"goal_patch_repair_failed:{reason}"
|
||||
return True, output, metadata, "ok"
|
||||
@@ -697,8 +724,10 @@ def run(job_path: str) -> int:
|
||||
try:
|
||||
validate_write_output(safe_local)
|
||||
except ValueError as first_error:
|
||||
stage(job_path, "local-worker", "running", "Repairing invalid patch output contract", job.get("local_provider", ""), local_model)
|
||||
emit(goal_id, "H2-tool", "running", "Local worker is repairing patch output contract", {"reason": str(first_error), "max_attempts": patch_repair_attempts()})
|
||||
repair_model = patch_repair_models(local_model)[0] if patch_repair_models(local_model) else local_model
|
||||
repair_provider = provider_for_model(repair_model)
|
||||
stage(job_path, "local-worker", "running", "Repairing invalid patch output contract", repair_provider, repair_model)
|
||||
emit(goal_id, "H2-tool", "running", "Worker is repairing patch output contract", {"reason": str(first_error), "max_attempts": patch_repair_attempts(), "repair_provider": repair_provider, "repair_model": repair_model})
|
||||
repaired, repaired_output, repaired_meta, repair_reason = repair_write_output(local_model, local_prompt, safe_local)
|
||||
local_meta = merge_usage(local_meta, repaired_meta)
|
||||
if not repaired:
|
||||
@@ -714,8 +743,9 @@ def run(job_path: str) -> int:
|
||||
except ValueError as repair_error:
|
||||
reason = f"goal_patch_repair_invalid:{repair_error}"
|
||||
if reason:
|
||||
stage(job_path, "local-worker", "error", reason, job.get("local_provider", ""), local_model)
|
||||
emit(goal_id, "H2-tool", "error", "Local worker violated patch output contract after repair", {"reason": reason})
|
||||
actual_repair_model = str(repaired_meta.get("repair_model") or repair_model)
|
||||
stage(job_path, "local-worker", "error", reason, provider_for_model(actual_repair_model), actual_repair_model)
|
||||
emit(goal_id, "H2-tool", "error", "Worker violated patch output contract after repair", {"reason": reason, "repair_provider": provider_for_model(actual_repair_model), "repair_model": actual_repair_model})
|
||||
raise ValueError(reason)
|
||||
stage(job_path, "local-worker", "pass", "Primary solution prepared", job.get("local_provider", ""), local_model)
|
||||
update_job(job_path, local_draft=safe_local, local_usage=local_meta)
|
||||
|
||||
@@ -57,6 +57,18 @@ class GoalPatchWorkflowTests(unittest.TestCase):
|
||||
self.assertEqual(reason, "goal_patch_missing")
|
||||
call.assert_not_called()
|
||||
|
||||
def test_repair_prefers_direct_cloud_over_gateway_and_local(self):
|
||||
repaired = "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@@ -1 +1 @@\n-a\n+b\n"
|
||||
with (
|
||||
patch.dict(os.environ, {"CASAN_GOAL_PATCH_REPAIR_ATTEMPTS": "1", "CASAN_GOAL_PATCH_REPAIR_MODELS": "openai:gpt-4o-mini,anthropic:claude,openai-compatible:auto/coding,ollama:ornith"}, clear=False),
|
||||
patch.object(ORCHESTRATOR, "call_model", return_value=(True, repaired, {}, "ok")) as call,
|
||||
):
|
||||
ok, _, usage, reason = ORCHESTRATOR.repair_write_output("ollama:ornith", "original", "plan")
|
||||
self.assertTrue(ok)
|
||||
self.assertEqual(reason, "ok")
|
||||
self.assertEqual(usage["repair_model"], "openai:gpt-4o-mini")
|
||||
self.assertEqual(call.call_args.args[0], "openai:gpt-4o-mini")
|
||||
|
||||
def test_patch_outside_workspace_is_denied(self):
|
||||
job = {"workspace": {"context_roots": ["apps/okr/frontend"]}}
|
||||
content = "diff --git a/package.json b/package.json\n--- a/package.json\n+++ b/package.json\n@@ -1 +1 @@\n-a\n+b\n"
|
||||
|
||||
Reference in New Issue
Block a user