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