Initial commit
This commit is contained in:
@@ -0,0 +1,300 @@
|
||||
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
|
||||
Reference in New Issue
Block a user