| from __future__ import annotations |
|
|
| import asyncio |
| import hashlib |
| import os |
| import time |
| from collections.abc import Mapping |
|
|
| from openai import AsyncOpenAI |
|
|
| from .config import WorkerConfig, WorkerPoolConfig |
| from .schemas import TaskRecord, WorkerResult |
|
|
|
|
| class WorkerCallError(RuntimeError): |
| pass |
|
|
|
|
| def _nested_extra(value, key: str): |
| extra = getattr(value, "model_extra", None) or {} |
| if isinstance(extra, Mapping): |
| return extra.get(key) |
| return None |
|
|
|
|
| class WorkerRunner: |
| def __init__(self, config: WorkerPoolConfig): |
| self.config = config |
| self._clients: dict[tuple[str, str], AsyncOpenAI] = {} |
| self._semaphore = asyncio.Semaphore(config.concurrency) |
|
|
| def _credentials(self, worker: WorkerConfig) -> tuple[str, str]: |
| base_url = worker.base_url or self.config.base_url |
| if worker.api_key: |
| key = worker.api_key |
| else: |
| env_name = worker.api_key_env or self.config.api_key_env |
| key = os.getenv(env_name, "") |
| if worker.provider == "openrouter" and not key: |
| raise WorkerCallError( |
| f"Missing {worker.api_key_env or self.config.api_key_env}; " |
| "put the OpenRouter key in the server environment or .env" |
| ) |
| return base_url, key or "EMPTY" |
|
|
| def _client(self, worker: WorkerConfig) -> AsyncOpenAI: |
| base_url, key = self._credentials(worker) |
| cache_key = (base_url, key) |
| if cache_key not in self._clients: |
| headers = {} |
| if os.getenv("OPENROUTER_HTTP_REFERER"): |
| headers["HTTP-Referer"] = os.environ["OPENROUTER_HTTP_REFERER"] |
| if os.getenv("OPENROUTER_APP_TITLE"): |
| headers["X-OpenRouter-Title"] = os.environ["OPENROUTER_APP_TITLE"] |
| self._clients[cache_key] = AsyncOpenAI( |
| base_url=base_url, |
| api_key=key, |
| timeout=self.config.timeout_seconds, |
| default_headers=headers or None, |
| ) |
| return self._clients[cache_key] |
|
|
| async def close(self) -> None: |
| for client in self._clients.values(): |
| await client.close() |
| self._clients.clear() |
|
|
| async def complete(self, worker: WorkerConfig, task: TaskRecord) -> WorkerResult: |
| if worker.provider == "mock": |
| return self._mock_complete(worker, task) |
| async with self._semaphore: |
| return await self._api_complete(worker, task) |
|
|
| def _mock_complete(self, worker: WorkerConfig, task: TaskRecord) -> WorkerResult: |
| started = time.perf_counter() |
| probability = ( |
| worker.mock_specialist_accuracy |
| if task.domain in worker.mock_domains |
| else worker.mock_general_accuracy |
| ) |
| digest = hashlib.sha256(f"{worker.id}:{task.task_id}".encode()).digest() |
| draw = int.from_bytes(digest[:8], "big") / float(2**64 - 1) |
| correct = draw < probability and task.reference_answer is not None |
| response = task.reference_answer if correct else f"incorrect mock answer from {worker.id}" |
| latency_ms = (time.perf_counter() - started) * 1000.0 + 1.0 |
| return WorkerResult( |
| worker_id=worker.id, |
| requested_model=worker.model, |
| served_model=worker.model, |
| response=response, |
| quality=0.0, |
| cost_usd=0.0, |
| latency_ms=latency_ms, |
| prompt_tokens=max(1, len(task.prompt.split())), |
| completion_tokens=max(1, len(response.split())), |
| generation_id=f"mock-{worker.id}-{task.task_id}", |
| ) |
|
|
| async def _api_complete(self, worker: WorkerConfig, task: TaskRecord) -> WorkerResult: |
| client = self._client(worker) |
| last_error: Exception | None = None |
| for attempt in range(self.config.max_retries): |
| started = time.perf_counter() |
| try: |
| response = await client.chat.completions.create( |
| model=worker.model, |
| messages=[ |
| {"role": "system", "content": worker.system_prompt}, |
| {"role": "user", "content": task.prompt}, |
| ], |
| temperature=worker.temperature, |
| max_tokens=worker.max_tokens, |
| extra_body=worker.extra_body or None, |
| ) |
| latency_ms = (time.perf_counter() - started) * 1000.0 |
| content = response.choices[0].message.content or "" |
| if not isinstance(content, str): |
| content = str(content) |
| usage = response.usage |
| cost = None |
| if usage is not None: |
| cost = getattr(usage, "cost", None) |
| if cost is None: |
| cost = _nested_extra(usage, "cost") |
| return WorkerResult( |
| worker_id=worker.id, |
| requested_model=worker.model, |
| served_model=response.model, |
| response=content, |
| quality=0.0, |
| cost_usd=float(cost) if cost is not None else None, |
| latency_ms=latency_ms, |
| prompt_tokens=getattr(usage, "prompt_tokens", None), |
| completion_tokens=getattr(usage, "completion_tokens", None), |
| generation_id=response.id, |
| ) |
| except Exception as exc: |
| last_error = exc |
| if attempt + 1 < self.config.max_retries: |
| await asyncio.sleep(min(2**attempt, 8)) |
| raise WorkerCallError( |
| f"Worker {worker.id} ({worker.model}) failed after " |
| f"{self.config.max_retries} attempts: {last_error}" |
| ) |
|
|
| async def judge(self, task: TaskRecord, candidate: str) -> float: |
| from .rewards import grader_prompt, parse_judge_score |
|
|
| model = self.config.reward.judge_model |
| if not model: |
| raise WorkerCallError( |
| f"Task {task.task_id} uses llm_judge but reward.judge_model is not configured" |
| ) |
| judge_worker = WorkerConfig( |
| id="__judge__", |
| provider="openrouter", |
| model=model, |
| system_prompt=self.config.reward.judge_system_prompt, |
| temperature=0.0, |
| max_tokens=32, |
| ) |
| judge_task = task.model_copy(update={"prompt": grader_prompt(task, candidate)}) |
| result = await self.complete(judge_worker, judge_task) |
| return parse_judge_score(result.response) |
|
|