fugu-lite / src /fugu_lite /providers.py
tahsinsoyak's picture
Upload private Fugu-Lite V1 snapshot
88e15cd verified
Raw
History Blame
6.69 kB
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)