from __future__ import annotations import json import math from dataclasses import asdict from pathlib import Path from typing import Any, Iterable from .models import Patch, RolloutTrace from .scoring.pre_score import PRE_SCORE_VERSION, PreScoreResult from .storage import ( atomic_write_json, load_json, package_hash, read_jsonl, sha256_file, sha256_text, ) CACHE_SCHEMA_VERSION = 1 CACHE_PRODUCER = "dynamic-compile-fast" SCORE_ARTIFACTS = ( "pre_scores.jsonl", "agentrm_scores.jsonl", "effective_scores.jsonl", ) def fingerprint(value: Any) -> str: payload = json.dumps( value, ensure_ascii=False, sort_keys=True, separators=(",", ":") ) return sha256_text(payload) def trace_fingerprint(traces: Iterable[RolloutTrace]) -> str: return fingerprint([asdict(trace) for trace in traces]) def _artifact_hash(path: Path) -> str: if path.is_file(): return sha256_file(path) if path.is_dir(): return package_hash(path) raise OSError(f"cache artifact does not exist: {path}") def write_manifest( root: Path, name: str, *, stage: str, input_hash: str, config: dict[str, Any], artifacts: Iterable[str], ) -> None: artifact_names = tuple(artifacts) hashes = {item: _artifact_hash(root / item) for item in artifact_names} atomic_write_json( root / name, { "schema_version": CACHE_SCHEMA_VERSION, "producer": CACHE_PRODUCER, "stage": stage, "status": "complete", "input_hash": input_hash, "config_hash": fingerprint(config), "artifacts": hashes, }, ) def valid_manifest( root: Path, name: str, *, stage: str, input_hash: str, config: dict[str, Any], artifacts: Iterable[str], ) -> bool: try: value = load_json(root / name) expected = tuple(artifacts) if not isinstance(value, dict): return False if value.get("schema_version") != CACHE_SCHEMA_VERSION: return False if value.get("producer") != CACHE_PRODUCER: return False if value.get("stage") != stage or value.get("status") != "complete": return False if value.get("input_hash") != input_hash: return False if value.get("config_hash") != fingerprint(config): return False hashes = value.get("artifacts") if not isinstance(hashes, dict) or set(hashes) != set(expected): return False return all( isinstance(hashes[item], str) and hashes[item] == _artifact_hash(root / item) for item in expected ) except (OSError, TypeError, ValueError): return False def load_score_cache( traces: list[RolloutTrace], score_dir: Path, *, input_hash: str, config: dict[str, Any], ) -> tuple[dict[str, float], list[dict[str, Any]]] | None: if not valid_manifest( score_dir, ".score-cache.json", stage="score", input_hash=input_hash, config=config, artifacts=SCORE_ARTIFACTS, ): return None try: pre_rows = read_jsonl(score_dir / "pre_scores.jsonl") agentrm_rows = read_jsonl(score_dir / "agentrm_scores.jsonl") score_rows = read_jsonl(score_dir / "effective_scores.jsonl") trace_by_id = {trace.trace_id: trace for trace in traces} if len(trace_by_id) != len(traces): return None pre_by_id: dict[str, PreScoreResult] = {} for row in pre_rows: result = PreScoreResult.from_dict(row) if result.trace_id in pre_by_id: return None if result.scoring_version != PRE_SCORE_VERSION: return None if result.route not in {"agentrm", "fixed_score"}: return None if result.route == "fixed_score": if result.score is None or not math.isfinite(result.score): return None pre_by_id[result.trace_id] = result if set(pre_by_id) != set(trace_by_id): return None expected_agentrm = { trace.key.as_tuple() for trace in traces if pre_by_id[trace.trace_id].route == "agentrm" } actual_agentrm: set[tuple[str, str, str]] = set() for row in agentrm_rows: key = ( str(row["task_name"]), str(row["compile_type"]), str(row["test_name"]), ) score = float(row["score"]) if key in actual_agentrm or not math.isfinite(score): return None if int(row["n_tokens"]) < 0: return None actual_agentrm.add(key) if actual_agentrm != expected_agentrm: return None by_id: dict[str, dict[str, Any]] = {} scores: dict[str, float] = {} for row in score_rows: trace_id = str(row["trace_id"]) score = float(row["effective_score"]) if trace_id in by_id or not math.isfinite(score): return None trace = trace_by_id.get(trace_id) if trace is None or ( str(row["task_name"]), str(row["compile_type"]), str(row["test_name"]), ) != trace.key.as_tuple(): return None expected_source = ( "agentrm" if pre_by_id[trace_id].route == "agentrm" else "pre_score" ) if row.get("score_source") != expected_source: return None by_id[trace_id] = row scores[trace_id] = score if set(by_id) != set(trace_by_id): return None return scores, [by_id[trace.trace_id] for trace in traces] except (KeyError, OSError, TypeError, ValueError): return None def load_maps_cache( output_dir: Path, selected: list[str], *, input_hash: str, config: dict[str, Any], ) -> list[dict[str, Any]] | None: if not valid_manifest( output_dir, ".maps-cache.json", stage="maps", input_hash=input_hash, config=config, artifacts=("maps.json",), ): return None try: value = load_json(output_dir / "maps.json") if not isinstance(value, list) or not all(isinstance(item, dict) for item in value): return None by_id = {str(item.get("trace_id", "")): item for item in value} if len(by_id) != len(value) or set(by_id) != set(selected): return None if any( not isinstance(item.get("patterns"), list) or not all(isinstance(pattern, dict) for pattern in item["patterns"]) for item in value ): return None return [by_id[trace_id] for trace_id in selected] except (OSError, TypeError, ValueError): return None def load_reduction_cache( output_dir: Path, *, input_hash: str, config: dict[str, Any], ) -> dict[str, Any] | None: if not valid_manifest( output_dir, ".reduction-cache.json", stage="reduction", input_hash=input_hash, config=config, artifacts=("reduction.json",), ): return None try: value = load_json(output_dir / "reduction.json") if not isinstance(value, dict): return None for key in ("successful_pattern", "failure_pattern", "selected_gap"): item = value.get(key) if not isinstance(item, dict) or not str(item.get("description", "")).strip(): return None return value except (OSError, TypeError, ValueError): return None def load_candidate_cache( output_dir: Path, candidate_name: str, skill_hash: str, *, input_hash: str, config: dict[str, Any], ) -> Path | None: candidate_relative = f"candidate-skill/{candidate_name}" if not valid_manifest( output_dir, ".candidate-cache.json", stage="candidate", input_hash=input_hash, config=config, artifacts=("patch.json", candidate_relative), ): return None candidate = output_dir / candidate_relative try: value = load_json(output_dir / "patch.json") if not isinstance(value, dict) or value.get("skill_hash") != skill_hash: return None patches = value.get("patches") if not isinstance(patches, list) or not 1 <= len(patches) <= 2: return None parsed = [Patch.from_dict(item) for item in patches if isinstance(item, dict)] if len(parsed) != len(patches) or any(patch.skill_hash != skill_hash for patch in parsed): return None if [patch.role for patch in parsed] != ["promote_success", "mitigate_failure"][:len(parsed)]: return None candidate_skill = candidate / "SKILL.md" if not candidate_skill.is_file(): return None if value.get("candidate_skill_hash") != sha256_file(candidate_skill): return None return candidate except (OSError, TypeError, ValueError): return None