Initial commit
This commit is contained in:
@@ -0,0 +1,191 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from typing import Any, Protocol
|
||||
|
||||
import requests
|
||||
|
||||
from .identity import composite_id
|
||||
|
||||
|
||||
DEFAULT_MAX_LENGTH = 8192
|
||||
DEFAULT_RM_API_URL = "http://127.0.0.1:28080"
|
||||
DEFAULT_TIMEOUT = 300.0
|
||||
DEFAULT_CONCURRENCY = 8
|
||||
DEFAULT_BATCH_SIZE = 32
|
||||
RETRY_ATTEMPTS = 3
|
||||
|
||||
|
||||
class AgentRMBackend(Protocol):
|
||||
def score(self, requests: list[dict[str, Any]]) -> list[dict[str, Any]]: ...
|
||||
|
||||
|
||||
class HttpAgentRMBackend:
|
||||
"""AgentRM backend backed by the remote ``/score_batch`` API."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_url: str | None = None,
|
||||
*,
|
||||
max_length: int = DEFAULT_MAX_LENGTH,
|
||||
timeout: float = DEFAULT_TIMEOUT,
|
||||
concurrency: int = DEFAULT_CONCURRENCY,
|
||||
batch_size: int = DEFAULT_BATCH_SIZE,
|
||||
) -> None:
|
||||
self.api_url = (
|
||||
api_url or os.environ.get("RM_API_URL", DEFAULT_RM_API_URL)
|
||||
).rstrip("/")
|
||||
self.max_length = max_length
|
||||
self.timeout = timeout
|
||||
self.concurrency = concurrency
|
||||
self.batch_size = batch_size
|
||||
if not self.api_url:
|
||||
raise ValueError("AgentRM API URL cannot be empty")
|
||||
if max_length <= 0 or timeout <= 0 or concurrency <= 0 or batch_size <= 0:
|
||||
raise ValueError(
|
||||
"AgentRM max_length, timeout, concurrency, and batch_size must be positive"
|
||||
)
|
||||
|
||||
def _post_batch(self, batch: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
payload = {
|
||||
"states": [request["state"] for request in batch],
|
||||
"max_length": self.max_length,
|
||||
}
|
||||
last_error: BaseException | None = None
|
||||
for attempt in range(RETRY_ATTEMPTS):
|
||||
try:
|
||||
response = requests.post(
|
||||
f"{self.api_url}/score_batch",
|
||||
json=payload,
|
||||
timeout=self.timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
body = response.json()
|
||||
scores = body.get("scores") if isinstance(body, dict) else None
|
||||
if not isinstance(scores, list) or len(scores) != len(batch):
|
||||
count = len(scores) if isinstance(scores, list) else "invalid"
|
||||
raise ValueError(
|
||||
f"AgentRM returned {count} scores for {len(batch)} states"
|
||||
)
|
||||
if any(not isinstance(score, dict) for score in scores):
|
||||
raise ValueError("AgentRM returned a non-object score item")
|
||||
if any(
|
||||
"score" not in score or "n_tokens" not in score
|
||||
for score in scores
|
||||
):
|
||||
raise ValueError("AgentRM returned a score item with missing fields")
|
||||
return [
|
||||
{
|
||||
"task_name": request["task_name"],
|
||||
"compile_type": request["compile_type"],
|
||||
"test_name": request["test_name"],
|
||||
**score,
|
||||
}
|
||||
for request, score in zip(batch, scores)
|
||||
]
|
||||
except (requests.RequestException, ValueError) as exc:
|
||||
last_error = exc
|
||||
if attempt + 1 < RETRY_ATTEMPTS:
|
||||
time.sleep(2**attempt)
|
||||
assert last_error is not None
|
||||
raise RuntimeError(
|
||||
f"AgentRM request failed after {RETRY_ATTEMPTS} attempts: {last_error}"
|
||||
)
|
||||
|
||||
def score(self, requests_to_score: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
if not requests_to_score:
|
||||
return []
|
||||
batches = [
|
||||
requests_to_score[index : index + self.batch_size]
|
||||
for index in range(0, len(requests_to_score), self.batch_size)
|
||||
]
|
||||
ordered: list[list[dict[str, Any]] | None] = [None] * len(batches)
|
||||
with ThreadPoolExecutor(max_workers=self.concurrency) as pool:
|
||||
futures = {
|
||||
pool.submit(self._post_batch, batch): index
|
||||
for index, batch in enumerate(batches)
|
||||
}
|
||||
for future in as_completed(futures):
|
||||
ordered[futures[future]] = future.result()
|
||||
return [row for batch in ordered if batch is not None for row in batch]
|
||||
|
||||
|
||||
class AgentRM:
|
||||
"""Score AgentRM requests through a validated, replaceable backend."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
backend: AgentRMBackend | None = None,
|
||||
*,
|
||||
api_url: str | None = None,
|
||||
max_length: int = DEFAULT_MAX_LENGTH,
|
||||
timeout: float = DEFAULT_TIMEOUT,
|
||||
concurrency: int = DEFAULT_CONCURRENCY,
|
||||
batch_size: int = DEFAULT_BATCH_SIZE,
|
||||
) -> None:
|
||||
self.backend = (
|
||||
backend
|
||||
if backend is not None
|
||||
else HttpAgentRMBackend(
|
||||
api_url,
|
||||
max_length=max_length,
|
||||
timeout=timeout,
|
||||
concurrency=concurrency,
|
||||
batch_size=batch_size,
|
||||
)
|
||||
)
|
||||
|
||||
def score_requests(
|
||||
self, requests_to_score: list[dict[str, Any]]
|
||||
) -> list[dict[str, Any]]:
|
||||
request_keys = [composite_id(request) for request in requests_to_score]
|
||||
if len(set(request_keys)) != len(request_keys):
|
||||
raise ValueError("AgentRM requests contain duplicate identities")
|
||||
|
||||
responses = self.backend.score(requests_to_score)
|
||||
response_by_key: dict[tuple[str, str, str], dict[str, Any]] = {}
|
||||
for response in responses:
|
||||
key = composite_id(response)
|
||||
if key in response_by_key:
|
||||
raise ValueError(f"AgentRM returned duplicate score identity: {key}")
|
||||
response_by_key[key] = response
|
||||
|
||||
requested = set(request_keys)
|
||||
unexpected = set(response_by_key) - requested
|
||||
missing = requested - set(response_by_key)
|
||||
if unexpected:
|
||||
raise ValueError(
|
||||
f"AgentRM returned unexpected score identities: {sorted(unexpected)}"
|
||||
)
|
||||
if missing:
|
||||
raise ValueError(f"AgentRM returned incomplete scores: {sorted(missing)}")
|
||||
|
||||
result = []
|
||||
for request, key in zip(requests_to_score, request_keys):
|
||||
response = response_by_key[key]
|
||||
try:
|
||||
score = float(response["score"])
|
||||
except (KeyError, TypeError, ValueError) as exc:
|
||||
raise ValueError(f"AgentRM returned an invalid score for {key}") from exc
|
||||
if not math.isfinite(score):
|
||||
raise ValueError(f"AgentRM returned a non-finite score for {key}")
|
||||
try:
|
||||
n_tokens = int(response["n_tokens"])
|
||||
except (KeyError, TypeError, ValueError) as exc:
|
||||
raise ValueError(
|
||||
f"AgentRM returned an invalid n_tokens for {key}"
|
||||
) from exc
|
||||
if n_tokens < 0:
|
||||
raise ValueError(f"AgentRM returned a negative n_tokens for {key}")
|
||||
row = {
|
||||
"task_name": str(request["task_name"]),
|
||||
"compile_type": str(request["compile_type"]),
|
||||
"test_name": str(request["test_name"]),
|
||||
"score": score,
|
||||
"n_tokens": n_tokens,
|
||||
}
|
||||
result.append(row)
|
||||
return result
|
||||
@@ -0,0 +1,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..models import TraceKey
|
||||
|
||||
|
||||
ScoreKey = tuple[str, str, str]
|
||||
|
||||
|
||||
def composite_id(row: dict[str, Any]) -> ScoreKey:
|
||||
return TraceKey.from_record(row).as_tuple()
|
||||
return result
|
||||
@@ -0,0 +1,254 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import math
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from dataclasses import asdict, dataclass, replace
|
||||
from typing import Any, Protocol
|
||||
|
||||
from ..models import RolloutTrace
|
||||
from ..optimization.analyzer import SemanticClient
|
||||
|
||||
|
||||
RELEVANCE_MODEL = "opencode/deepseek-v4-pro"
|
||||
RELEVANCE_CONFIDENCE_THRESHOLD = 0.80
|
||||
TIMEOUT_SCORE = 0.0
|
||||
IRRELEVANT_SCORE = 0.1
|
||||
MAX_EFFICIENCY_PENALTY = 0.08
|
||||
TIME_COST_WEIGHT = 0.80
|
||||
TOOL_COST_WEIGHT = 0.20
|
||||
EFFICIENCY_WINSOR_QUANTILE = 0.90
|
||||
PRE_SCORE_VERSION = 3
|
||||
|
||||
|
||||
class RelevanceClient(Protocol):
|
||||
def json(self, system: str, user: str, attempts: int = 3) -> dict[str, Any]: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PreScoreResult:
|
||||
trace_id: str
|
||||
route: str
|
||||
score: float | None
|
||||
reason: str
|
||||
relevance_label: str | None = None
|
||||
relevance_confidence: float | None = None
|
||||
efficiency_cost: float = 0.0
|
||||
efficiency_penalty: float = 0.0
|
||||
scoring_version: int = PRE_SCORE_VERSION
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return asdict(self)
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, value: dict[str, Any]) -> "PreScoreResult":
|
||||
return cls(
|
||||
trace_id=str(value["trace_id"]),
|
||||
route=str(value["route"]),
|
||||
score=float(value["score"]) if value.get("score") is not None else None,
|
||||
reason=str(value["reason"]),
|
||||
relevance_label=(
|
||||
str(value["relevance_label"])
|
||||
if value.get("relevance_label") is not None
|
||||
else None
|
||||
),
|
||||
relevance_confidence=(
|
||||
float(value["relevance_confidence"])
|
||||
if value.get("relevance_confidence") is not None
|
||||
else None
|
||||
),
|
||||
efficiency_cost=float(value.get("efficiency_cost", 0.0)),
|
||||
efficiency_penalty=float(value.get("efficiency_penalty", 0.0)),
|
||||
scoring_version=int(value.get("scoring_version", 1)),
|
||||
)
|
||||
|
||||
|
||||
def adjust_agent_score(score: float, pre_score: PreScoreResult) -> float:
|
||||
"""Apply the batch-relative efficiency penalty to an AgentRM quality score."""
|
||||
if pre_score.route != "agentrm":
|
||||
return score
|
||||
return max(IRRELEVANT_SCORE, score - pre_score.efficiency_penalty)
|
||||
|
||||
|
||||
def _quantile(values: list[float], quantile: float) -> float:
|
||||
ordered = sorted(values)
|
||||
if len(ordered) == 1:
|
||||
return ordered[0]
|
||||
position = (len(ordered) - 1) * quantile
|
||||
lower = math.floor(position)
|
||||
upper = math.ceil(position)
|
||||
if lower == upper:
|
||||
return ordered[lower]
|
||||
fraction = position - lower
|
||||
return ordered[lower] + fraction * (ordered[upper] - ordered[lower])
|
||||
|
||||
|
||||
def _magnitude_costs(
|
||||
values: dict[str, float],
|
||||
*,
|
||||
transform=lambda value: value,
|
||||
power: float = 1.0,
|
||||
) -> dict[str, float]:
|
||||
if len(values) < 2:
|
||||
return {trace_id: 0.0 for trace_id in values}
|
||||
transformed = {trace_id: float(transform(value)) for trace_id, value in values.items()}
|
||||
floor = min(transformed.values())
|
||||
ceiling = _quantile(list(transformed.values()), EFFICIENCY_WINSOR_QUANTILE)
|
||||
if ceiling <= floor:
|
||||
return {trace_id: 0.0 for trace_id in values}
|
||||
scale = ceiling - floor
|
||||
return {
|
||||
trace_id: min(1.0, max(0.0, (value - floor) / scale)) ** power
|
||||
for trace_id, value in transformed.items()
|
||||
}
|
||||
|
||||
|
||||
class RelevanceJudge:
|
||||
def __init__(
|
||||
self,
|
||||
client: RelevanceClient | None = None,
|
||||
model: str = RELEVANCE_MODEL,
|
||||
):
|
||||
self.model = model
|
||||
self.client = client or SemanticClient(model)
|
||||
|
||||
def judge(self, task_prompt: str, final_output: str) -> tuple[str, float]:
|
||||
result = self.client.json(
|
||||
"Judge only whether an agent's last output is relevant to its task. Return JSON only.",
|
||||
f"""Classify the last agent output as relevant or irrelevant to the task.
|
||||
|
||||
An output is relevant if it attempts, plans, discusses, or reports work on the requested task, even when it is wrong, incomplete, brief, malformed, or lacks a final answer. Mark it irrelevant only when it clearly addresses a materially different task or topic. Do not judge correctness, completeness, or answer quality.
|
||||
|
||||
Return exactly {{"label":"relevant"|"irrelevant","confidence":number}} where confidence is between 0 and 1.
|
||||
|
||||
Input:
|
||||
{json.dumps({"task_prompt": task_prompt, "last_agent_output": final_output}, ensure_ascii=False)}""",
|
||||
)
|
||||
if not isinstance(result, dict) or set(result) != {"label", "confidence"}:
|
||||
raise ValueError("relevance response must contain exactly label and confidence")
|
||||
label = result.get("label")
|
||||
confidence = result.get("confidence")
|
||||
if label not in {"relevant", "irrelevant"}:
|
||||
raise ValueError("relevance label must be relevant or irrelevant")
|
||||
if isinstance(confidence, bool) or not isinstance(confidence, (int, float)):
|
||||
raise ValueError("relevance confidence must be a number")
|
||||
confidence = float(confidence)
|
||||
if not math.isfinite(confidence) or not 0.0 <= confidence <= 1.0:
|
||||
raise ValueError("relevance confidence must be between 0 and 1")
|
||||
return label, confidence
|
||||
|
||||
|
||||
class PreScorer:
|
||||
def __init__(
|
||||
self,
|
||||
judge: RelevanceJudge,
|
||||
max_parallel: int = 3,
|
||||
confidence_threshold: float = RELEVANCE_CONFIDENCE_THRESHOLD,
|
||||
):
|
||||
self.judge = judge
|
||||
self.max_parallel = max(1, max_parallel)
|
||||
self.confidence_threshold = confidence_threshold
|
||||
|
||||
@staticmethod
|
||||
def _first_user_message(trace: RolloutTrace) -> str:
|
||||
return next(
|
||||
(
|
||||
str(message.get("content", "")).strip()
|
||||
for message in trace.state
|
||||
if isinstance(message, dict)
|
||||
and message.get("role") == "user"
|
||||
and str(message.get("content", "")).strip()
|
||||
),
|
||||
"",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _last_assistant_message(trace: RolloutTrace) -> str:
|
||||
return next(
|
||||
(
|
||||
str(message.get("content", "")).strip()
|
||||
for message in reversed(trace.state)
|
||||
if isinstance(message, dict)
|
||||
and message.get("role") == "assistant"
|
||||
and str(message.get("content", "")).strip()
|
||||
),
|
||||
"",
|
||||
)
|
||||
|
||||
def score_one(self, trace: RolloutTrace) -> PreScoreResult:
|
||||
if trace.timed_out:
|
||||
return PreScoreResult(trace.trace_id, "fixed_score", TIMEOUT_SCORE, "hard_timeout")
|
||||
|
||||
final_output = self._last_assistant_message(trace)
|
||||
if not final_output:
|
||||
return PreScoreResult(trace.trace_id, "agentrm", None, "no_assistant_output")
|
||||
|
||||
task_prompt = self._first_user_message(trace)
|
||||
try:
|
||||
label, confidence = self.judge.judge(task_prompt, final_output)
|
||||
except Exception:
|
||||
return PreScoreResult(trace.trace_id, "agentrm", None, "judge_failed")
|
||||
|
||||
if label == "irrelevant" and confidence >= self.confidence_threshold:
|
||||
return PreScoreResult(
|
||||
trace.trace_id,
|
||||
"fixed_score",
|
||||
IRRELEVANT_SCORE,
|
||||
"strongly_irrelevant",
|
||||
label,
|
||||
confidence,
|
||||
)
|
||||
reason = "low_confidence" if label == "irrelevant" else "relevant"
|
||||
return PreScoreResult(
|
||||
trace.trace_id, "agentrm", None, reason, label, confidence
|
||||
)
|
||||
|
||||
def score_all(self, traces: list[RolloutTrace]) -> list[PreScoreResult]:
|
||||
results: list[PreScoreResult | None] = [None] * len(traces)
|
||||
with ThreadPoolExecutor(max_workers=self.max_parallel) as pool:
|
||||
futures = {
|
||||
pool.submit(self.score_one, trace): index
|
||||
for index, trace in enumerate(traces)
|
||||
}
|
||||
for future in as_completed(futures):
|
||||
results[futures[future]] = future.result()
|
||||
completed = [result for result in results if result is not None]
|
||||
by_id = {trace.trace_id: trace for trace in traces}
|
||||
eligible = {
|
||||
result.trace_id for result in completed if result.route == "agentrm"
|
||||
}
|
||||
durations = {
|
||||
trace_id: float(by_id[trace_id].metadata["agent_execution_seconds"])
|
||||
for trace_id in eligible
|
||||
if isinstance(
|
||||
by_id[trace_id].metadata.get("agent_execution_seconds"), (int, float)
|
||||
)
|
||||
}
|
||||
tool_calls = {
|
||||
trace_id: float(by_id[trace_id].metadata["tool_calls"])
|
||||
for trace_id in eligible
|
||||
if isinstance(by_id[trace_id].metadata.get("tool_calls"), (int, float))
|
||||
}
|
||||
duration_cost = _magnitude_costs(durations, power=2.0)
|
||||
tool_cost = _magnitude_costs(tool_calls, transform=math.log1p)
|
||||
adjusted = []
|
||||
for result in completed:
|
||||
components = []
|
||||
if result.trace_id in duration_cost:
|
||||
components.append((TIME_COST_WEIGHT, duration_cost[result.trace_id]))
|
||||
if result.trace_id in tool_cost:
|
||||
components.append((TOOL_COST_WEIGHT, tool_cost[result.trace_id]))
|
||||
total_weight = sum(weight for weight, _ in components)
|
||||
cost = (
|
||||
sum(weight * value for weight, value in components) / total_weight
|
||||
if total_weight
|
||||
else 0.0
|
||||
)
|
||||
adjusted.append(
|
||||
replace(
|
||||
result,
|
||||
efficiency_cost=cost,
|
||||
efficiency_penalty=MAX_EFFICIENCY_PENALTY * cost,
|
||||
)
|
||||
)
|
||||
return adjusted
|
||||
@@ -0,0 +1,75 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from ..models import RolloutTrace
|
||||
from ..storage import atomic_write_jsonl
|
||||
from .agentrm import AgentRM
|
||||
from .pre_score import PreScorer, adjust_agent_score
|
||||
from .identity import composite_id
|
||||
|
||||
|
||||
class TraceScorer:
|
||||
"""Resolve pre-scores and AgentRM scores into one effective score per trace."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pre_scorer: PreScorer,
|
||||
agentrm: AgentRM | None = None,
|
||||
):
|
||||
self.pre_scorer = pre_scorer
|
||||
self.agentrm = agentrm if agentrm is not None else AgentRM()
|
||||
|
||||
def score_all(
|
||||
self,
|
||||
traces: list[RolloutTrace],
|
||||
final_score_dir: Path,
|
||||
agentrm_dir: Path | None = None,
|
||||
) -> dict[str, float]:
|
||||
final_score_dir.mkdir(parents=True, exist_ok=True)
|
||||
agentrm_dir = agentrm_dir or final_score_dir
|
||||
agentrm_dir.mkdir(parents=True, exist_ok=True)
|
||||
results = self.pre_scorer.score_all(traces)
|
||||
pre_scores = {result.trace_id: result for result in results}
|
||||
if set(pre_scores) != {trace.trace_id for trace in traces}:
|
||||
raise ValueError("pre-scorer returned incomplete or duplicate results")
|
||||
atomic_write_jsonl(
|
||||
final_score_dir / "pre_scores.jsonl",
|
||||
[pre_scores[trace.trace_id].to_dict() for trace in traces],
|
||||
)
|
||||
agentrm_traces = [
|
||||
trace for trace in traces if pre_scores[trace.trace_id].route == "agentrm"
|
||||
]
|
||||
requests = [trace.agentrm_request() for trace in agentrm_traces]
|
||||
local_scores = agentrm_dir / "agentrm_scores.jsonl"
|
||||
found = self.agentrm.score_requests(requests)
|
||||
atomic_write_jsonl(local_scores, found)
|
||||
|
||||
by_key = {composite_id(row): float(row["score"]) for row in found}
|
||||
resolved = []
|
||||
scores: dict[str, float] = {}
|
||||
for trace in traces:
|
||||
pre_score = pre_scores[trace.trace_id]
|
||||
if pre_score.route == "fixed_score":
|
||||
if pre_score.score is None:
|
||||
raise ValueError(f"fixed pre-score is missing for {trace.trace_id}")
|
||||
score = pre_score.score
|
||||
source = "pre_score"
|
||||
else:
|
||||
score = adjust_agent_score(by_key[trace.key.as_tuple()], pre_score)
|
||||
source = "agentrm"
|
||||
scores[trace.trace_id] = score
|
||||
resolved.append({
|
||||
"task_name": trace.task_name,
|
||||
"compile_type": trace.compile_type,
|
||||
"test_name": trace.test_name,
|
||||
"trace_id": trace.trace_id,
|
||||
"effective_score": score,
|
||||
"score_source": source,
|
||||
"pre_score_reason": pre_score.reason,
|
||||
"agentrm_score": by_key.get(trace.key.as_tuple()),
|
||||
"efficiency_cost": pre_score.efficiency_cost,
|
||||
"efficiency_penalty": pre_score.efficiency_penalty,
|
||||
})
|
||||
atomic_write_jsonl(final_score_dir / "effective_scores.jsonl", resolved)
|
||||
return scores
|
||||
@@ -0,0 +1,121 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable
|
||||
|
||||
|
||||
MAX_CHARS = 2000
|
||||
LAST_MSG_BUDGET = 4000
|
||||
HEAD_RATIO = 0.5
|
||||
|
||||
|
||||
def compact(content: str, budget: int) -> str:
|
||||
if len(content) <= budget:
|
||||
return content
|
||||
head = int(budget * HEAD_RATIO)
|
||||
tail = budget - head - 50
|
||||
return (
|
||||
content[:head]
|
||||
+ f"\n...[truncated {len(content)-head-tail} chars]...\n"
|
||||
+ content[-tail:]
|
||||
)
|
||||
|
||||
|
||||
def _unwrap_text_content(value: str) -> str:
|
||||
prefix = "@{type=text; text="
|
||||
if value.startswith(prefix) and value.endswith("}"):
|
||||
return value[len(prefix) : -1]
|
||||
return value
|
||||
|
||||
|
||||
def _tool_result_text(event: dict[str, Any]) -> str:
|
||||
parts: list[str] = []
|
||||
content_items = event.get("content")
|
||||
if isinstance(content_items, list):
|
||||
for item in content_items:
|
||||
if not isinstance(item, dict) or item.get("type") != "content":
|
||||
continue
|
||||
value = item.get("content")
|
||||
if isinstance(value, str) and value:
|
||||
parts.append(_unwrap_text_content(value))
|
||||
elif (
|
||||
isinstance(value, dict)
|
||||
and value.get("type") == "text"
|
||||
and isinstance(value.get("text"), str)
|
||||
and value["text"]
|
||||
):
|
||||
parts.append(value["text"])
|
||||
return "\n".join(parts) if parts else "(no output)"
|
||||
|
||||
|
||||
def acp_events_to_state(events: Iterable[dict[str, Any]]) -> list[dict[str, str]]:
|
||||
merged: list[dict[str, str]] = []
|
||||
for event in events:
|
||||
event_type = event.get("type")
|
||||
role: str | None = None
|
||||
content: str | None = None
|
||||
if event_type == "user_message":
|
||||
role, content = "user", event.get("text")
|
||||
elif event_type == "agent_thought":
|
||||
role, content = "assistant", event.get("text")
|
||||
elif event_type == "agent_message":
|
||||
role, content = "assistant", event.get("text")
|
||||
if content == "":
|
||||
continue
|
||||
elif event_type == "tool_call":
|
||||
role = "user"
|
||||
kind = event.get("kind", "other")
|
||||
status = event.get("status", "unknown")
|
||||
content = f"Tool result ({kind}; {status}):\n{_tool_result_text(event)}"
|
||||
elif event_type == "agent_timeout":
|
||||
continue
|
||||
else:
|
||||
continue
|
||||
|
||||
if not isinstance(content, str):
|
||||
raise ValueError(f"{event_type} event has a non-string text/content value")
|
||||
if merged and merged[-1]["role"] == role:
|
||||
merged[-1]["content"] += "\n" + content
|
||||
else:
|
||||
merged.append({"role": role, "content": content})
|
||||
|
||||
if not merged:
|
||||
raise ValueError("trajectory produced an empty state")
|
||||
if merged[0]["role"] != "user":
|
||||
raise ValueError("state must start with a user message")
|
||||
|
||||
seen: set[str] = set()
|
||||
last_index = len(merged) - 1
|
||||
for index, message in enumerate(merged):
|
||||
budget = LAST_MSG_BUDGET if index == last_index else MAX_CHARS
|
||||
content = compact(message["content"], budget)
|
||||
if len(content) > 200:
|
||||
digest = hashlib.md5(content.encode("utf-8")).hexdigest()
|
||||
if digest in seen:
|
||||
content = "[same as previous tool result]"
|
||||
else:
|
||||
seen.add(digest)
|
||||
message["content"] = content
|
||||
return merged
|
||||
|
||||
|
||||
def read_acp_events(path: Path) -> list[dict[str, Any]]:
|
||||
events: list[dict[str, Any]] = []
|
||||
with path.open("r", encoding="utf-8") as handle:
|
||||
for line_number, line in enumerate(handle, 1):
|
||||
if not line.strip():
|
||||
continue
|
||||
try:
|
||||
event = json.loads(line)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(f"invalid JSON on line {line_number}: {exc.msg}") from exc
|
||||
if not isinstance(event, dict):
|
||||
raise ValueError(f"line {line_number} is not a JSON object")
|
||||
events.append(event)
|
||||
return events
|
||||
|
||||
|
||||
def read_acp_state(path: Path) -> list[dict[str, str]]:
|
||||
return acp_events_to_state(read_acp_events(path))
|
||||
Reference in New Issue
Block a user