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: # noqa: BLE001 - provider SDKs expose many transient types 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)