Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
hint_generator.py418 linesDownload Raw Back to composer_replication
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