301 lines
9.1 KiB
Python
301 lines
9.1 KiB
Python
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
|