Codeseys/composer-replication-framework
0
1"""hint_generator.py — Template-based hint generator (v0.1 starter).2 3Composer 2.5 inserts text hints at error-turn sites:4 "Reminder: Available tools are: …" (when a tool-call refs a non-existent tool)5 "Reminder: tool arguments must be valid JSON" (on JSONDecodeError)6 ... etc.7 8This module provides a registry of hint templates keyed by error_kind. The9data collator (in trl_path/data_collator.py) calls dispatch(error_kind, ctx)10to get the hint text to splice into ctx_teacher.11 12v0.2 will replace these templates with an LLM-driven hint generator (likely13Sonnet 4.6 or Opus 4.7 via OpenRouter) for cases where templates are too rigid14(style violations, wasteful explanations).15"""16 17from __future__ import annotations18 19from collections.abc import Callable20from typing import TypedDict21 22 23class HintContext(TypedDict, total=False):24 """Per-error context the hint generator can use."""25 error_kind: str # e.g. "tool_not_found", "json_decode", "type_error"26 error_message: str # raw error from the env27 available_tools: list[str] # for tool_not_found28 tool_name: str # the failing tool, if known29 tool_schema: dict # the schema, if known30 intent: str # student's apparent intent, if extractable31 32 33# ---------------------------------------------------------------------------34# Hint templates35# ---------------------------------------------------------------------------36 37def hint_tool_not_found(ctx: HintContext) -> str:38 tools = ctx.get("available_tools", [])39 if tools:40 tool_list = ", ".join(f"`{t}`" for t in tools)41 return f"Reminder: Available tools are: {tool_list}. Please use one of these."42 return "Reminder: the tool you tried to call does not exist. Use only available tools."43 44 45def hint_json_decode(ctx: HintContext) -> str:46 return (47 "Reminder: tool arguments must be valid JSON. Common mistakes: "48 "single quotes (use double), trailing commas, unescaped newlines in strings."49 )50 51 52def hint_type_error(ctx: HintContext) -> str:53 name = ctx.get("tool_name")54 schema = ctx.get("tool_schema")55 if name and schema:56 return (57 f"Reminder: `{name}` expects arguments matching this schema:\n"58 f" {schema}\n"59 "Re-issue the call with arguments matching the schema."60 )61 return "Reminder: tool arguments do not match the expected types. Check the schema."62 63 64def hint_runtime_error(ctx: HintContext) -> str:65 msg = ctx.get("error_message", "an exception")66 return (67 f"Reminder: the previous tool call raised {msg}. "68 "Reconsider the inputs or read the relevant code first to understand state."69 )70 71 72def hint_repeated_failure(ctx: HintContext) -> str:73 """Triggered when the same kind of error happens 3+ times in a row."""74 return (75 "Reminder: this approach has failed multiple times. "76 "Step back and consider an alternative approach: read more files, "77 "search for similar patterns elsewhere, or break the task down differently."78 )79 80 81# ---------------------------------------------------------------------------82# Registry83# ---------------------------------------------------------------------------84 85HINT_TEMPLATES: dict[str, Callable[[HintContext], str]] = {86 "tool_not_found": hint_tool_not_found,87 "json_decode": hint_json_decode,88 "type_error": hint_type_error,89 "runtime_error": hint_runtime_error,90 "repeated_failure": hint_repeated_failure,91}92 93 94def dispatch(error_kind: str, ctx: HintContext | None = None) -> str | None:95 """Generate a hint for the given error_kind. Returns None if unknown."""96 fn = HINT_TEMPLATES.get(error_kind)97 if fn is None:98 return None99 return fn(ctx or {})100 101 102def register(error_kind: str, fn: Callable[[HintContext], str]) -> None:103 """Add a custom hint template."""104 HINT_TEMPLATES[error_kind] = fn105 106 107# ===========================================================================108# Layered HintGenerator architecture (ADR-009)109# ===========================================================================110#111# Composer 2.5 inserts a natural-language hint at each error turn; the112# hint-conditioned forward becomes the SDPO teacher. HOW Cursor generates the113# hint is unstated in every Cursor artifact (both blogs + the Composer 2 tech114# report, arXiv:2603.24477 — confirmed absent in research/10). So this is our115# design problem. The cited papers bracket the answer: OPSD conditions the116# teacher on ground-truth; SDPO generalizes to environment feedback and the117# "successful sibling rollout as implicit feedback" trick.118#119# We implement a layered generator, tried cheapest-first:120# 1. TemplateHintGenerator — the registry above (free, deterministic;121# covers tool-error classes). The first layer.122# 2. RawErrorHintGenerator — wrap the raw env/tool error text as the hint123# (free; covers any error with a message but unmatched by a template).124# 3. LLMJudgeHintGenerator — an LLM produces a <=2-sentence corrective hint125# (cost ~$0.0005/site; covers style/communication/effort sites templates126# can't). Cached on disk; optional; OFF unless a client is provided.127# 4. (sibling-bootstrap) — RL-rollout-path only; not a HintContext-driven128# layer (needs sibling rollouts), exposed as a flag for the trainer to use.129#130# All layers satisfy the HintGenerator Protocol and compose via131# CompositeHintGenerator, whose .as_collator_hook() returns a callable matching132# the collator's existing `hint_generator: Callable[[str, dict], str | None]`133# hook — ZERO collator change.134 135from typing import Protocol, runtime_checkable136 137 138@runtime_checkable139class HintGenerator(Protocol):140 """A hint source. Returns hint text for an error context, or None to defer141 to the next layer."""142 143 def generate(self, error_kind: str, error_meta: dict) -> str | None: ...144 145 146class TemplateHintGenerator:147 """Layer 1: the existing template registry. Free, deterministic.148 149 Preserves the exact behavior of the module-level `dispatch()` so existing150 callers and tests see no change.151 """152 153 def generate(self, error_kind: str, error_meta: dict) -> str | None:154 # `dispatch` reads HintContext keys; error_meta IS that context dict155 # plus the kind. Merge so templates that read `error_kind` still work.156 ctx: HintContext = dict(error_meta) # type: ignore[assignment]157 ctx.setdefault("error_kind", error_kind)158 return dispatch(error_kind, ctx)159 160 161class RawErrorHintGenerator:162 """Layer 2: use the raw env/tool error text itself as the hint.163 164 Covers any error site that carries a message but isn't matched by a165 template. Free. SDPO's "environment feedback as the conditioning signal"166 (arXiv:2601.20802) — the rawest form of that.167 """168 169 def __init__(self, max_chars: int = 500) -> None:170 self.max_chars = max_chars171 172 def generate(self, error_kind: str, error_meta: dict) -> str | None:173 msg = error_meta.get("error_message") or error_meta.get("error") or ""174 msg = str(msg).strip()175 if not msg:176 return None177 truncated = msg[: self.max_chars]178 return f"Reminder: the previous action produced this error:\n{truncated}\nReconsider and retry."179 180 181# ---------------------------------------------------------------------------182# Error-kind routing (ADR-012 finding #2)183# ---------------------------------------------------------------------------184#185# The default composite is template -> raw-error -> judge. The raw-error layer186# fires for ANY kind carrying a message — including style/communication/effort187# sites, which are EXACTLY what the LLM judge exists to cover. So we route:188# tool/runtime error kinds may use the raw-error layer; style/communication/189# effort kinds skip it and fall through to the judge.190 191# Error kinds that genuinely describe a tool/runtime failure whose raw text is a192# useful, self-contained hint. The explicit registry-template kinds are included193# so behavior is unchanged for them.194_TOOL_RUNTIME_KINDS: frozenset[str] = frozenset({195 "tool_not_found",196 "json_decode",197 "type_error",198 "runtime_error",199 "repeated_failure",200})201 202# Substrings marking a kind as tool/runtime-ish even if not explicitly listed203# (keeps generic "*_error"/"*_exception" sites flowing through raw-error, which204# is where their raw text belongs).205_TOOL_RUNTIME_MARKERS: tuple[str, ...] = (206 "error", "exception", "fail", "decode", "timeout", "traceback",207 "exit_code", "nonzero", "syntax", "import", "assertion", "tool",208 "runtime", "crash", "exec",209)210 211# Substrings marking a kind as a style/communication/effort site — the judge's212# domain. These take precedence: a kind matching one of these skips raw-error.213_STYLE_KINDS_MARKERS: tuple[str, ...] = (214 "style", "communic", "verbose", "effort", "concise", "tone",215 "format", "wordy", "rambl", "explanation", "etiquette", "clarity",216)217 218 219def is_tool_runtime_kind(error_kind: str) -> bool:220 """True if `error_kind` is a tool/runtime failure that the raw-error layer221 may serve. Style/communication/effort kinds return False (-> judge)."""222 k = (error_kind or "").lower()223 if any(m in k for m in _STYLE_KINDS_MARKERS):224 return False225 if k in _TOOL_RUNTIME_KINDS:226 return True227 return any(m in k for m in _TOOL_RUNTIME_MARKERS)228 229 230class RoutingHintGenerator:231 """Wraps an inner layer (the raw-error layer) and only lets it fire for232 tool/runtime error kinds. For style/communication/effort kinds it returns233 None so the composite falls through to the judge — the layer those sites234 were always meant to reach (ADR-012 finding #2).235 """236 237 def __init__(self, inner: HintGenerator, route=is_tool_runtime_kind) -> None:238 self.inner = inner239 self.route = route240 241 def generate(self, error_kind: str, error_meta: dict) -> str | None:242 if not self.route(error_kind):243 return None244 return self.inner.generate(error_kind, error_meta)245 246 247class LLMJudgeHintGenerator:248 """Layer 3: an LLM produces a short corrective hint.249 250 Covers style/communication/effort sites that templates can't. Optional and251 OFF unless a `complete` callable is provided. Results are cached on disk252 keyed on a hash of the error context (so repeated identical sites cost253 nothing after the first).254 255 `complete(prompt: str) -> str` is an injected text-completion callable256 (e.g. an OpenRouter chat wrapper). Kept abstract so this module has no hard257 network dependency and is unit-testable with a stub.258 """259 260 PROMPT_TEMPLATE = (261 "An autonomous coding agent made a mistake at one step of a trajectory. "262 "Write a SHORT (<=2 sentences) corrective hint that, if the agent had "263 "seen it, would steer it to the right behavior for THIS step only. Do "264 "not solve the whole task; just correct the local mistake.\n\n"265 "Error kind: {error_kind}\n"266 "Error / context:\n{error_message}\n\n"267 "Corrective hint:"268 )269 270 # Bump when PROMPT_TEMPLATE or the underlying judge model changes so stale271 # cached hints are invalidated rather than silently reused.272 _CACHE_VERSION = 2273 274 # Hard cap on a generated hint. The judge is asked for <=2 sentences but275 # nothing enforced it (cross-family review 2026-05-29) — a runaway judge276 # could emit a full solution / prompt-leak / megabyte of text straight into277 # the SDPO teacher conditioning. Clamp defensively.278 _MAX_HINT_CHARS = 600279 280 def __init__(281 self,282 complete: Callable[[str], str] | None = None,283 *,284 cache_dir: str | None = None,285 ) -> None:286 self.complete = complete287 self._cache_dir = cache_dir288 self._mem_cache: dict[str, str] = {}289 290 def _cache_key(self, error_kind: str, error_meta: dict) -> str:291 import hashlib292 import json293 import re294 295 # Strip volatile object reprs (e.g. "<Exception at 0x7f8b...>") so the296 # key is stable across runs/restarts. Cross-family review 2026-05-29:297 # `default=str` on raw Exception/context objects embedded a memory298 # address in the key, guaranteeing a 0% cross-process cache-hit rate and299 # unbounded judge cost. Also version the key so prompt/model changes300 # invalidate stale hints rather than serving them.301 blob = json.dumps(302 {"v": self._CACHE_VERSION, "k": error_kind, "m": error_meta},303 sort_keys=True, default=str,304 )305 blob = re.sub(r"0x[0-9a-fA-F]+", "0xADDR", blob)306 blob = re.sub(r"\bat 0xADDR\b", "", blob)307 return hashlib.sha256(blob.encode("utf-8")).hexdigest()[:32]308 309 def _disk_get(self, key: str) -> str | None:310 if not self._cache_dir:311 return None312 from pathlib import Path313 314 p = Path(self._cache_dir) / f"{key}.txt"315 return p.read_text(encoding="utf-8") if p.exists() else None316 317 def _disk_put(self, key: str, value: str) -> None:318 if not self._cache_dir:319 return320 import os321 from pathlib import Path322 323 d = Path(self._cache_dir)324 d.mkdir(parents=True, exist_ok=True)325 # Atomic write: concurrent DDP workers writing the same key would326 # otherwise interleave and corrupt the file (cross-family review).327 tmp = d / f"{key}.txt.{os.getpid()}.tmp"328 tmp.write_text(value, encoding="utf-8")329 os.replace(tmp, d / f"{key}.txt")330 331 def generate(self, error_kind: str, error_meta: dict) -> str | None:332 if self.complete is None:333 return None # judge disabled — defer334 key = self._cache_key(error_kind, error_meta)335 if key in self._mem_cache:336 return self._mem_cache[key]337 cached = self._disk_get(key)338 if cached is not None:339 self._mem_cache[key] = cached340 return cached341 prompt = self.PROMPT_TEMPLATE.format(342 error_kind=error_kind,343 error_message=str(error_meta.get("error_message")344 or error_meta.get("error") or "(no message)")[:1000],345 )346 hint = self.complete(prompt).strip()347 if not hint:348 return None349 # Clamp to a sane length so a runaway judge can't inject a full solution350 # or megabyte blob into the SDPO teacher conditioning (cross-family review).351 if len(hint) > self._MAX_HINT_CHARS:352 hint = hint[: self._MAX_HINT_CHARS].rstrip() + "…"353 self._mem_cache[key] = hint354 self._disk_put(key, hint)355 return hint356 357 358class CompositeHintGenerator:359 """Tries each layer in order, returning the first non-None hint.360 361 Order is cost-ascending: templates (free) -> raw error (free) -> LLM judge362 (paid, optional). The first layer to produce a hint wins, so the common363 tool-error case never reaches the LLM.364 """365 366 def __init__(self, layers: list[HintGenerator]) -> None:367 self.layers = layers368 369 def generate(self, error_kind: str, error_meta: dict) -> str | None:370 for layer in self.layers:371 hint = layer.generate(error_kind, error_meta)372 if hint is not None:373 return hint374 return None375 376 def as_collator_hook(self) -> Callable[[str, dict], str | None]:377 """Return a callable matching CollatorConfig.hint_generator's signature378 (error_kind, error_meta) -> str | None. ZERO collator change."""379 return self.generate380 381 382def default_composite(383 *,384 llm_complete: Callable[[str], str] | None = None,385 cache_dir: str | None = None,386 enable_raw_error: bool = True,387) -> CompositeHintGenerator:388 """Build the recommended layered generator: templates -> raw-error -> judge.389 390 The raw-error layer is wrapped in a RoutingHintGenerator so it only fires for391 tool/runtime error kinds; style/communication/effort kinds skip it and fall392 through to the LLM judge (ADR-012 finding #2). The LLM-judge layer is393 included only when `llm_complete` is provided.394 """395 layers: list[HintGenerator] = [TemplateHintGenerator()]396 if enable_raw_error:397 layers.append(RoutingHintGenerator(RawErrorHintGenerator()))398 if llm_complete is not None:399 layers.append(LLMJudgeHintGenerator(llm_complete, cache_dir=cache_dir))400 return CompositeHintGenerator(layers)401 402 403__all__ = [404 "dispatch",405 "register",406 "HintContext",407 "HINT_TEMPLATES",408 # Layered architecture (ADR-009)409 "HintGenerator",410 "TemplateHintGenerator",411 "RawErrorHintGenerator",412 "RoutingHintGenerator",413 "is_tool_runtime_kind",414 "LLMJudgeHintGenerator",415 "CompositeHintGenerator",416 "default_composite",417]418 