from __future__ import annotations import json import re import time import uuid from pathlib import Path, PurePosixPath from typing import Any from calibration_agent.supervisor.pipeline import PipelineClient, parse_prompt_jsonl from calibration_agent.supervisor.remote import RemoteRunner from calibration_agent.supervisor.settings import SupervisorSettings from calibration_agent.supervisor.store import JsonlStore def _now() -> str: return time.strftime("%Y-%m-%dT%H:%M:%S%z") def _slug(text: str, limit: int = 36) -> str: value = re.sub(r"[^A-Za-z0-9._-]+", "-", text.strip())[:limit].strip("-") return value or "requirement" class RequirementStore: def __init__(self, settings: SupervisorSettings): self.settings = settings self.root = settings.requirements_dir def create(self, payload: dict[str, Any]) -> dict[str, Any]: req_id = f"{time.strftime('%Y%m%d_%H%M%S')}_{uuid.uuid4().hex[:8]}" record = { "id": req_id, "created_at": _now(), "updated_at": _now(), "status": "active", "requirement": payload["requirement"], "success_criteria": payload.get("success_criteria", ""), "max_iterations": int(payload.get("max_iterations", 3)), "artifact": payload.get("artifact", "thinker_outputs_latest"), "sample_limit": int(payload.get("sample_limit", 6)), "prompt_keys": payload.get("prompt_keys", []), "propose_prompt_keys": payload.get("propose_prompt_keys", []), "target_behavior": payload.get("target_behavior", ""), "default_stage": payload.get("default_stage", "status"), "active_prompt_path": "", "active_data_root": "", "active_iteration": 0, "current_activity": { "phase": "created", "message": "Requirement created.", "updated_at": _now(), }, "iterations": [], } self._write(record) return record def list(self) -> list[dict[str, Any]]: records = [] for path in sorted(self.root.glob("*.json")): try: records.append(json.loads(path.read_text(encoding="utf-8"))) except Exception: continue return sorted(records, key=lambda item: item.get("updated_at", ""), reverse=True) def load(self, req_id: str) -> dict[str, Any]: path = self._path(req_id) if not path.exists(): raise KeyError(f"Requirement not found: {req_id}") return json.loads(path.read_text(encoding="utf-8")) def append_iteration(self, req_id: str, iteration: dict[str, Any]) -> dict[str, Any]: record = self.load(req_id) record.setdefault("iterations", []).append(iteration) record["active_iteration"] = iteration["iteration_index"] if iteration.get("active_prompt_path"): record["active_prompt_path"] = iteration["active_prompt_path"] if iteration.get("active_data_root"): record["active_data_root"] = iteration["active_data_root"] if iteration.get("is_solved"): record["status"] = "solved" record["updated_at"] = _now() self._write(record) return record def set_activity( self, req_id: str, phase: str, message: str, details: dict[str, Any] | None = None, ) -> dict[str, Any]: record = self.load(req_id) record["current_activity"] = { "phase": phase, "message": message, "details": details or {}, "updated_at": _now(), } record["updated_at"] = _now() self._write(record) return record["current_activity"] def rollback(self, req_id: str, iteration_index: int) -> dict[str, Any]: record = self.load(req_id) target = None for item in record.get("iterations", []): if int(item.get("iteration_index", -1)) == iteration_index: target = item break if target is None: raise KeyError(f"Iteration not found: {iteration_index}") record["active_iteration"] = iteration_index record["active_prompt_path"] = target.get("active_prompt_path", "") record["active_data_root"] = target.get("active_data_root", "") record["status"] = "active" record["updated_at"] = _now() self._write(record) return record def next_iteration_index(self, record: dict[str, Any]) -> int: return len(record.get("iterations", [])) + 1 def _path(self, req_id: str) -> Path: return self.root / f"{req_id}.json" def _write(self, record: dict[str, Any]) -> None: self._path(record["id"]).write_text( json.dumps(record, ensure_ascii=False, indent=2), encoding="utf-8", ) class ExperimentManager: def __init__( self, settings: SupervisorSettings, runner: RemoteRunner, pipeline: PipelineClient, audit_store: JsonlStore, ): self.settings = settings self.runner = runner self.pipeline = pipeline self.audit_store = audit_store def snapshot( self, requirement: dict[str, Any], iteration_index: int, sample: dict[str, Any] | None, ) -> dict[str, Any]: base = ( self.settings.requirements_dir / requirement["id"] / f"iter_{iteration_index:03d}" / "snapshot" ) base.mkdir(parents=True, exist_ok=True) config = self.pipeline.get_config() prompts = self.runner.read_text( requirement.get("active_prompt_path") or self.settings.prompt_path ) status = self.pipeline.status() (base / "config.yaml").write_text(config["raw"], encoding="utf-8") (base / "prompt.jsonl").write_text(prompts, encoding="utf-8") (base / "status.json").write_text( json.dumps(status, ensure_ascii=False, indent=2), encoding="utf-8", ) if sample is not None: (base / "sample.json").write_text( json.dumps(sample, ensure_ascii=False, indent=2), encoding="utf-8", ) return { "local_dir": str(base), "prompt_source": requirement.get("active_prompt_path") or self.settings.prompt_path, "config_source": self.settings.config_path, "data_root": config["values"].get("DATA_ROOT", ""), "status_counts": status.get("counts", {}), "status_data_root": status.get("data_root", ""), } def get_prompt_from_active_version( self, requirement: dict[str, Any], key: str, ) -> dict[str, Any]: path = requirement.get("active_prompt_path") or self.settings.prompt_path raw = self.runner.read_text(path) for row in parse_prompt_jsonl(raw): if row.get("key") == key: row["_source_path"] = path return row raise KeyError(f"Prompt key not found in {path}: {key}") def create_prompt_version( self, requirement: dict[str, Any], iteration_index: int, updates: list[dict[str, str]], execute: bool, ) -> dict[str, Any]: if not updates: return {"skipped": True, "reason": "No prompt update selected."} source_path = requirement.get("active_prompt_path") or self.settings.prompt_path raw = self.runner.read_text(source_path) rows = parse_prompt_jsonl(raw) changed = [] for update in updates: key = update["key"] text = update["text"] for row in rows: if row.get("key") == key: row["text"] = text changed.append(key) break else: rows.append({"key": key, "text": text, "description": ""}) changed.append(key) new_text = "\n".join(json.dumps(row, ensure_ascii=False) for row in rows) + "\n" local_dir = ( self.settings.requirements_dir / requirement["id"] / f"iter_{iteration_index:03d}" / "versions" ) local_dir.mkdir(parents=True, exist_ok=True) local_path = local_dir / "prompt.jsonl" local_path.write_text(new_text, encoding="utf-8") remote_path = ( f"{self.settings.pipeline_dir}/archive/supervisor_runs/" f"{requirement['id']}/iter_{iteration_index:03d}/prompt.jsonl" ) result = { "changed_keys": changed, "source_path": source_path, "local_path": str(local_path), "remote_path": remote_path, "dry_run": not execute, } if execute: write_result, _ = self.runner.write_text(remote_path, new_text, backup_dir=None) result["write_result"] = write_result.to_dict() self.audit_store.append("requirement.prompt_version", result) return result def data_root_for_iteration( self, requirement: dict[str, Any], iteration_index: int, ) -> str: current = self.pipeline.get_config()["values"].get("DATA_ROOT", "") if not current: return "" parent = PurePosixPath(str(current)).parent req_slug = _slug(requirement.get("requirement", requirement["id"])) return str(parent / "supervisor_runs" / req_slug / f"iter_{iteration_index:03d}")