parhamkhoshsolat's picture
Provider pool + summary cache/cap
8eb87a5 verified
Raw
History Blame Contribute Delete
4.77 kB
from __future__ import annotations
import os
from typing import Optional, Protocol
class LLMClient(Protocol):
def complete(self, system: str, user: str) -> str: ...
class FakeLLMClient:
def __init__(self, response: str = "Answer [S1]."):
self.response = response
self.calls: list[tuple[str, str]] = []
def complete(self, system: str, user: str) -> str:
self.calls.append((system, user))
return self.response
class RaisingClient:
"""Test helper that always raises, to exercise failover."""
def __init__(self, exc: Optional[Exception] = None):
self.exc = exc or RuntimeError("provider down")
def complete(self, system: str, user: str) -> str:
raise self.exc
class OpenAICompatClient:
"""Works for any OpenAI-compatible endpoint (Groq, Cerebras). Lazy import."""
def __init__(self, api_key: str, base_url: str, model: str, temperature: float = 0.2):
from openai import OpenAI
self._client = OpenAI(api_key=api_key, base_url=base_url)
self.model = model
self.temperature = temperature
def complete(self, system: str, user: str) -> str:
resp = self._client.chat.completions.create(
model=self.model,
messages=[{"role": "system", "content": system}, {"role": "user", "content": user}],
temperature=self.temperature,
)
return resp.choices[0].message.content or ""
class FailoverClient:
"""Best-provider-first failover with a cooldown: a provider that fails (daily cap,
bad key, outage) is skipped for `cooldown` seconds instead of being retried on
every request — so a capped provider costs one failed call, not added latency on
every answer for the rest of the day."""
def __init__(self, clients: list[LLMClient], cooldown: float = 600.0):
if not clients:
raise ValueError("FailoverClient needs at least one client")
self.clients = clients
self.cooldown = cooldown
self._down: dict[int, float] = {}
def complete(self, system: str, user: str) -> str:
import time as _time
now = _time.time()
last: Optional[Exception] = None
for i, c in enumerate(self.clients):
if self._down.get(i, 0.0) > now:
continue
try:
out = c.complete(system, user)
self._down.pop(i, None)
return out
except Exception as exc:
self._down[i] = now + self.cooldown
last = exc
if last is None:
# Everyone is cooling down — desperation pass: maybe a quota just reset.
for i, c in enumerate(self.clients):
try:
out = c.complete(system, user)
self._down.pop(i, None)
return out
except Exception as exc:
last = exc
raise last
_PROVIDERS = (
# (env key, base_url, model env, default model) — quality order, ~70B-class first.
("GROQ_API_KEY", "https://api.groq.com/openai/v1", "GROQ_MODEL", "llama-3.3-70b-versatile"),
("CEREBRAS_API_KEY", "https://api.cerebras.ai/v1", "CEREBRAS_MODEL", "qwen-3-32b"),
("SAMBANOVA_API_KEY", "https://api.sambanova.ai/v1", "SAMBANOVA_MODEL", "Meta-Llama-3.3-70B-Instruct"),
("GEMINI_API_KEY", "https://generativelanguage.googleapis.com/v1beta/openai/",
"GEMINI_MODEL", "gemini-2.0-flash"),
("OPENROUTER_API_KEY", "https://openrouter.ai/api/v1",
"OPENROUTER_MODEL", "meta-llama/llama-3.3-70b-instruct:free"),
)
def default_failover_from_env() -> FailoverClient:
"""Build the provider pool from whatever keys exist: each free tier has its own
independent daily quota, so every extra key adds a full day's worth of capacity.
Optionally STUDYHUB_HF_LLM_MODEL adds the HF Inference router (PRO credits)."""
clients: list[LLMClient] = []
for env_key, base, model_env, default_model in _PROVIDERS:
if os.getenv(env_key):
clients.append(OpenAICompatClient(os.environ[env_key], base,
os.getenv(model_env, default_model)))
if os.getenv("STUDYHUB_HF_LLM_MODEL") and os.getenv("HF_TOKEN"):
clients.append(OpenAICompatClient(os.environ["HF_TOKEN"],
"https://router.huggingface.co/v1",
os.environ["STUDYHUB_HF_LLM_MODEL"]))
if not clients:
raise RuntimeError("Set at least one of GROQ_API_KEY, CEREBRAS_API_KEY, SAMBANOVA_API_KEY, "
"GEMINI_API_KEY, OPENROUTER_API_KEY (or STUDYHUB_HF_LLM_MODEL with HF_TOKEN).")
return FailoverClient(clients)