Download src/peppa/engine.py from ChatterjeeLab/PepPA: direct link, hf CLI and curl.
- Browser
- Download file 6.74 kB
-
https://huggingface.co/ChatterjeeLab/PepPA/resolve/main/src/peppa/engine.py
- Command line
-
hf download hf://ChatterjeeLab/PepPA/src/peppa/engine.py
-
curl -L -o engine.py https://huggingface.co/ChatterjeeLab/PepPA/resolve/main/src/peppa/engine.py
6.74 kB
| """Bounded execution of explicit controller decisions against a tool registry.""" | |
| from __future__ import annotations | |
| import json | |
| import math | |
| import copy | |
| import subprocess | |
| import tempfile | |
| from pathlib import Path | |
| from typing import Callable | |
| from .schema import Decision, DesignSpec, ToolResult, canonical | |
| from .trace import Trace, replay | |
| class Tool: | |
| def __init__(self, name: str, function: Callable, cost: dict[str,float], | |
| description: str, argument_schema: dict | None = None): | |
| if any(not math.isfinite(v) or v < 0 for v in cost.values()): | |
| raise ValueError("tool costs must be nonnegative") | |
| self.name,self.function,self.cost,self.description=name,function,cost,description | |
| self.argument_schema=argument_schema or {} | |
| def __call__(self, arguments, state): | |
| if self.argument_schema.get("type"): | |
| from jsonschema import validate | |
| validate(arguments,self.argument_schema) | |
| return ToolResult.model_validate(self.function(arguments,copy.deepcopy(state))) | |
| class CommandTool(Tool): | |
| """Execute a preconfigured JSON-in/JSON-out worker in its own environment. | |
| argv is supplied by the operator's registry, never by a controller decision. | |
| Workers receive a request file and an output file as their last arguments. | |
| """ | |
| def __init__(self, name, argv, cwd, cost, description, timeout=1800, argument_schema=None): | |
| self.argv=list(argv); self.cwd=Path(cwd).resolve(); self.timeout=timeout | |
| super().__init__(name,self.execute,cost,description,argument_schema) | |
| def execute(self, arguments, state): | |
| with tempfile.TemporaryDirectory(prefix="peppa-job-") as directory: | |
| request=Path(directory)/"request.json";output=Path(directory)/"result.json" | |
| request.write_text(canonical({"arguments":arguments,"state":state})) | |
| completed=subprocess.run(self.argv+[str(request),str(output)],cwd=self.cwd, | |
| capture_output=True,text=True,timeout=self.timeout,check=False) | |
| if completed.returncode: | |
| raise RuntimeError(f"worker exit {completed.returncode}: {completed.stderr[-2000:]}") | |
| if not output.exists(): | |
| raise RuntimeError("worker did not write its declared result") | |
| return json.loads(output.read_text()) | |
| class Engine: | |
| def __init__(self, spec: DesignSpec, registry: dict[str,Tool], trace_path: str | Path): | |
| self.registry=registry;self.trace=Trace(trace_path) | |
| if not self.trace.events: | |
| self.trace.append("initialize",spec.model_dump(mode="json")) | |
| self.state=replay(self.trace.path) | |
| if self.state["spec"] != spec.model_dump(mode="json"): | |
| raise ValueError("resume specification differs from original episode") | |
| def reserve(self, name: str, cost: dict[str,float]): | |
| limits=self.state["spec"]["budgets"] | |
| for resource,amount in cost.items(): | |
| if not math.isfinite(amount) or amount<0 or resource not in limits or self.state["spent"].get(resource,0)+amount>limits[resource]: | |
| raise ValueError(f"budget exceeded or undefined: {resource}") | |
| self.trace.append("reserve",{"tool":name,"cost":cost}) | |
| self.state=replay(self.trace.path) | |
| def apply(self, decision: Decision): | |
| if self.state["stopped"]: | |
| raise ValueError("episode already stopped") | |
| unknown=set(decision.evidence_ids)-set(self.state["evidence"]) | |
| if unknown: | |
| raise ValueError(f"unknown evidence IDs: {sorted(unknown)}") | |
| self.trace.append("decision",decision.model_dump(mode="json")) | |
| if decision.stop: | |
| self.trace.append("stop",{"summary":decision.decision_summary}) | |
| elif decision.tool=="revise_plan": | |
| # Scientific endpoints and budgets are immutable. Planning weights and | |
| # questions may change, with the visible evidence retained in the trace. | |
| allowed={"weights","motifs","questions","next_assays","hypotheses"} | |
| if set(decision.arguments)-allowed: | |
| raise ValueError("revision attempts to change a fixed scientific requirement") | |
| weights=decision.arguments.get("weights") | |
| if weights is not None: | |
| if not isinstance(weights,dict) or not weights or any(float(v)<0 for v in weights.values()) or sum(weights.values())<=0: | |
| raise ValueError("invalid objective weights") | |
| names={r["endpoint"] for r in self.state["spec"]["requirements"]} | |
| if set(weights)-names: | |
| raise ValueError("unknown objective name") | |
| self.reserve("revise_plan",{"tool_calls":1}) | |
| self.trace.append("revise",decision.arguments) | |
| else: | |
| if decision.tool not in self.registry: | |
| raise ValueError("tool is not registered") | |
| tool=self.registry[decision.tool] | |
| self.reserve(tool.name,tool.cost) | |
| try: | |
| result=tool(decision.arguments,self.state) | |
| known=set(self.state["candidates"])|{c.id for c in result.candidates} | |
| if any(m.candidate_id not in known for m in result.measurements): | |
| raise ValueError("measurement references an unknown candidate") | |
| self.trace.append("result",result.model_dump(mode="json")) | |
| except Exception as exc: | |
| self.trace.append("error",{"tool":tool.name,"type":type(exc).__name__,"message":str(exc)}) | |
| self.state=replay(self.trace.path) | |
| raise | |
| self.state=replay(self.trace.path) | |
| return self.state | |
| def run(self, controller, max_steps=40): | |
| for _ in range(max_steps): | |
| if self.state["stopped"]:break | |
| try: | |
| cost={"controller_calls":1} | |
| if getattr(controller,"token_reservation",0):cost["controller_tokens"]=controller.token_reservation | |
| self.reserve("controller",cost) | |
| except ValueError as exc: | |
| self.trace.append("stop",{"summary":str(exc)});self.state=replay(self.trace.path);break | |
| try: | |
| decision,metadata=controller.next(copy.deepcopy(self.state),self.registry) | |
| self.trace.append("controller",metadata) | |
| self.apply(decision) | |
| except Exception as exc: | |
| self.trace.append("feedback",{"error":str(exc)}) | |
| self.state=replay(self.trace.path) | |
| if not self.state["stopped"]: | |
| self.trace.append("stop",{"summary":"controller step limit reached"}) | |
| self.state=replay(self.trace.path) | |
| return self.state | |