from __future__ import annotations from typing import Any, Literal from pydantic import BaseModel, Field, model_validator class GraderSpec(BaseModel): type: Literal["exact", "contains", "numeric", "regex", "llm_judge"] = "exact" tolerance: float = 0.0 pattern: str | None = None rubric: str | None = None case_sensitive: bool = False class TaskRecord(BaseModel): task_id: str prompt: str domain: str = "general" reference_answer: str | None = None grader: GraderSpec = Field(default_factory=GraderSpec) split: Literal["train", "validation", "test"] | None = None tags: list[str] = Field(default_factory=list) metadata: dict[str, Any] = Field(default_factory=dict) class WorkerResult(BaseModel): worker_id: str requested_model: str served_model: str | None = None response: str quality: float cost_usd: float | None = None latency_ms: float prompt_tokens: int | None = None completion_tokens: int | None = None generation_id: str | None = None error: str | None = None utility: float = 0.0 class RewardRecord(BaseModel): task_id: str prompt: str domain: str split: Literal["train", "validation", "test"] | None = None tags: list[str] = Field(default_factory=list) metadata: dict[str, Any] = Field(default_factory=dict) repetitions: int = Field(default=1, ge=1) worker_ids: list[str] rewards: list[float] results: list[WorkerResult] = Field(default_factory=list) @model_validator(mode="after") def validate_vector(self) -> RewardRecord: if not self.worker_ids: raise ValueError("worker_ids cannot be empty") if len(self.worker_ids) != len(self.rewards): raise ValueError("worker_ids and rewards must have equal length") if len(set(self.worker_ids)) != len(self.worker_ids): raise ValueError("worker_ids must be unique") return self def route_text( prompt: str, domain: str = "general", tags: list[str] | None = None, worker_descriptions: dict[str, str] | None = None, ) -> str: """Create the router input. Do not include answers or worker outcomes.""" tag_text = ", ".join(tags or []) or "none" experts = "" if worker_descriptions: experts = "\nAvailable workers:\n" + "\n".join( f"- {worker_id}: {description}" for worker_id, description in worker_descriptions.items() ) return ( "Select the best worker for this task.\n" f"Domain: {domain}\n" f"Tags: {tag_text}{experts}\n" f"Task:\n{prompt}" )