90 lines
2.7 KiB
Python
90 lines
2.7 KiB
Python
from __future__ import annotations
|
|
|
|
from dataclasses import asdict, dataclass, field
|
|
from typing import Any, Mapping
|
|
|
|
|
|
@dataclass(frozen=True, order=True)
|
|
class TraceKey:
|
|
task_name: str
|
|
compile_type: str
|
|
test_name: str
|
|
|
|
@classmethod
|
|
def from_record(cls, value: Mapping[str, Any]) -> "TraceKey":
|
|
try:
|
|
return cls(
|
|
str(value["task_name"]),
|
|
str(value["compile_type"]),
|
|
str(value["test_name"]),
|
|
)
|
|
except KeyError as exc:
|
|
raise ValueError(f"missing trace identity field: {exc.args[0]}") from exc
|
|
|
|
def as_tuple(self) -> tuple[str, str, str]:
|
|
return self.task_name, self.compile_type, self.test_name
|
|
|
|
def __str__(self) -> str:
|
|
return "/".join(self.as_tuple())
|
|
|
|
|
|
@dataclass
|
|
class Patch:
|
|
edit_type: str
|
|
target_heading: str
|
|
old_text: str
|
|
new_text: str
|
|
evidence_ids: list[str]
|
|
evidence_type: str
|
|
confidence: str
|
|
reason: str
|
|
role: str = ""
|
|
skill_hash: str = ""
|
|
|
|
@classmethod
|
|
def from_dict(cls, value: dict[str, Any]) -> "Patch":
|
|
required = {
|
|
"edit_type", "target_heading", "old_text", "new_text", "evidence_ids",
|
|
"evidence_type", "confidence", "reason",
|
|
}
|
|
missing = sorted(required - value.keys())
|
|
if missing:
|
|
raise ValueError(f"patch missing fields: {', '.join(missing)}")
|
|
if not isinstance(value["evidence_ids"], list):
|
|
raise ValueError("patch evidence_ids must be a list")
|
|
if "role" in value and not isinstance(value["role"], str):
|
|
raise ValueError("patch role must be a string")
|
|
for key in required - {"evidence_ids"}:
|
|
if not isinstance(value[key], str):
|
|
raise ValueError(f"patch {key} must be a string")
|
|
return cls(**{key: value.get(key, "") for key in cls.__dataclass_fields__})
|
|
|
|
def to_dict(self) -> dict[str, Any]:
|
|
return asdict(self)
|
|
|
|
|
|
@dataclass
|
|
class RolloutTrace:
|
|
trace_id: str
|
|
task_name: str
|
|
compile_type: str
|
|
test_name: str
|
|
state: list[dict[str, Any]]
|
|
skill_invoked: bool
|
|
exit_code: int | None = None
|
|
timed_out: bool | None = None
|
|
metadata: dict[str, Any] = field(default_factory=dict)
|
|
|
|
@property
|
|
def key(self) -> TraceKey:
|
|
return TraceKey(self.task_name, self.compile_type, self.test_name)
|
|
|
|
def agentrm_request(self) -> dict[str, Any]:
|
|
"""Project a rich runtime trace onto AgentRM's stable input schema."""
|
|
return {
|
|
"state": self.state,
|
|
"task_name": self.task_name,
|
|
"compile_type": self.compile_type,
|
|
"test_name": self.test_name,
|
|
}
|