""" Multi-provider LLM Client for DS Theorist — rotating keys, OpenAI-compatible. Provider chain (updated 2026-04-03): 1. Groq — llama-3.3-70b-versatile (PRIMARY — GROQ_API_KEY or GROQ_KEY_1..8) 2. Cerebras — qwen-3-235b-a22b-instruct-2507 (FALLBACK — CEREBRAS_API_KEY or CEREBRAS_KEY_1..4) 3. Sarvam — sarvam-m (SECOND FALLBACK — SARVAM_API_KEY) 4. NVIDIA — meta/llama-3.3-70b-instruct (legacy keys) 5. Mistral — mistral-small-latest (legacy keys) 6. Inception — mercury-2 (legacy keys) Optimized for mathematical reasoning and formal proofs. """ import os import time import httpx # ── Provider definitions ───────────────────────────────────────────────────── def _load_keys(prefix: str, count: int) -> list[str]: return [k for k in [os.getenv(f"{prefix}_{i}") for i in range(1, count + 1)] if k] def _load_single_or_multi(single_env: str, multi_prefix: str, count: int) -> list[str]: """Load a single env var key first, then fall back to numbered keys.""" keys = [] single = os.getenv(single_env) if single: keys.append(single) keys += _load_keys(multi_prefix, count) # deduplicate while preserving order seen = set() result = [] for k in keys: if k not in seen: seen.add(k) result.append(k) return result PROVIDERS = [] # 1. Groq (PRIMARY — fast, free) _groq_keys = _load_single_or_multi("GROQ_API_KEY", "GROQ_KEY", 8) if _groq_keys: PROVIDERS.append({ "name": "Groq", "base": "https://api.groq.com/openai/v1/chat/completions", "model": "llama-3.3-70b-versatile", "model_fast": "llama-3.3-70b-versatile", "keys": _groq_keys, "auth": "Bearer", "timeout": 120.0, "max_tokens_cap": 4096, }) # 2. Cerebras (FALLBACK — best for math, qwen-3-235b) _cerebras_keys = _load_single_or_multi("CEREBRAS_API_KEY", "CEREBRAS_KEY", 4) if _cerebras_keys: PROVIDERS.append({ "name": "Cerebras", "base": "https://api.cerebras.ai/v1/chat/completions", "model": "qwen-3-235b-a22b-instruct-2507", "model_fast": "llama3.1-8b", "keys": _cerebras_keys, "auth": "Bearer", "timeout": 120.0, "max_tokens_cap": 4096, }) # 3. Sarvam (SECOND FALLBACK) _sarvam_keys = [k for k in [os.getenv("SARVAM_API_KEY")] if k] if _sarvam_keys: PROVIDERS.append({ "name": "Sarvam", "base": "https://api.sarvam.ai/v1/chat/completions", "model": "sarvam-m", "model_fast": "sarvam-m", "keys": _sarvam_keys, "auth": "Bearer", "timeout": 120.0, "max_tokens_cap": 4096, }) # 4. NVIDIA (legacy free credits) _nvidia_keys = _load_keys("NVAPI_KEY", 4) if _nvidia_keys: PROVIDERS.append({ "name": "NVIDIA", "base": "https://integrate.api.nvidia.com/v1/chat/completions", "model": "meta/llama-3.3-70b-instruct", "model_fast": "meta/llama-3.3-70b-instruct", "keys": _nvidia_keys, "auth": "Bearer", "timeout": 120.0, "max_tokens_cap": 4096, }) # 5. Mistral (legacy — solid quality) _mistral_keys = _load_keys("MISTRAL_KEY", 4) if _mistral_keys: PROVIDERS.append({ "name": "Mistral", "base": "https://api.mistral.ai/v1/chat/completions", "model": "mistral-small-latest", "model_fast": "mistral-small-latest", "keys": _mistral_keys, "auth": "Bearer", "timeout": 120.0, "max_tokens_cap": 4096, }) # 6. Inception (legacy — mercury-2) _inception_keys = _load_keys("INCEPTION_KEY", 8) if _inception_keys: PROVIDERS.append({ "name": "Inception", "base": "https://api.inceptionlabs.ai/v1/chat/completions", "model": "mercury-2", "model_fast": "mercury-2", "keys": _inception_keys, "auth": "Bearer", "timeout": 120.0, "max_tokens_cap": 4096, }) # ── Key rotation state ─────────────────────────────────────────────────────── _key_indices: dict[str, int] = {} def _next_key(provider: dict) -> str: name = provider["name"] idx = _key_indices.get(name, 0) key = provider["keys"][idx % len(provider["keys"])] _key_indices[name] = (idx + 1) % len(provider["keys"]) return key def _make_headers(provider: dict, key: str) -> dict: auth_type = provider.get("auth", "Bearer") if auth_type == "api-subscription-key": return {"api-subscription-key": key, "Content-Type": "application/json"} return {"Authorization": f"Bearer {key}", "Content-Type": "application/json"} # ── Main API call ──────────────────────────────────────────────────────────── def complete( messages: list, max_tokens: int = 4096, temperature: float = 0.72, fast: bool = False, ) -> str: """ Call LLM chat completion with automatic provider fallback. Optimized for mathematical reasoning and formal proofs. """ if not PROVIDERS: raise RuntimeError("LLM: no providers configured — set GROQ_API_KEY, CEREBRAS_API_KEY, or SARVAM_API_KEY in env") last_error = "no providers" for provider in PROVIDERS: model = provider["model_fast"] if fast else provider["model"] cap = provider.get("max_tokens_cap", 4096) capped_tokens = min(max_tokens, cap) if fast: capped_tokens = min(capped_tokens, 300) min_tok = provider.get("min_tokens", 1) capped_tokens = max(capped_tokens, min_tok) payload = { "model": model, "messages": messages, "max_tokens": capped_tokens, "temperature": temperature, "stream": False, } for _attempt in range(len(provider["keys"])): key = _next_key(provider) headers = _make_headers(provider, key) try: r = httpx.post( provider["base"], headers=headers, json=payload, timeout=provider.get("timeout", 120.0), ) if r.status_code == 429: last_error = f"{provider['name']} 429 rate-limited" time.sleep(5) continue if r.status_code in (401, 403): last_error = f"{provider['name']} auth error (key=...{key[-6:]})" continue if r.status_code == 402: last_error = f"{provider['name']} 402 credits exhausted" break r.raise_for_status() data = r.json() return data["choices"][0]["message"]["content"].strip() except httpx.HTTPStatusError as e: last_error = f"{provider['name']} HTTP {e.response.status_code}" if e.response.status_code in (400, 404, 402): break except Exception as e: last_error = f"{provider['name']}: {type(e).__name__}: {e}" break raise RuntimeError(f"LLM (DS Theorist): all providers failed. Last: {last_error}")