Spaces:
Running on Zero
Running on Zero
File size: 9,143 Bytes
f1ef7e2 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 | """
Provider-agnostic LLM client (OpenAI SDK-compatible for all providers).
Swap provider with a single string -- the rest of the pipeline never changes.
Used by extract.py via: chat_json(provider, system, user)
Provider notes
--------------
github : GitHub Models free tier -- GPT-4o-mini, zero cost, recommended default
openai : OpenAI direct -- costs credits, use for prod/comparison only
gemini : Google Gemini via OpenAI-compat endpoint -- free tier, 15 RPM daily cap
mistral : Mistral AI -- mistral-small is strong at structured JSON
cohere : Cohere -- command-r7b fastest (326ms), good at extraction
nvidia : NVIDIA NIM -- llama-3.1-8b, solid and fast
cerebras : Cerebras -- zai-glm-4.7 is a reasoning model (needs higher token budget)
cloudflare : Cloudflare Workers AI -- needs CLOUDFLARE_ID in .env
huggingface: HF Inference via featherless-ai -- slowest, cold-start ~14s
sambanova : SambaNova -- has DeepSeek-V3 + Llama-4, fast on warm runs
openrouter : OpenRouter -- free models available, quality varies
deepseek : DeepSeek -- needs paid balance (402 currently)
JSON mode support
-----------------
Providers that natively support response_format=json_object are flagged
json_mode=True. Others receive a JSON-enforcement suffix in the system prompt
and the response is extracted via regex fallback if needed.
"""
from __future__ import annotations
import os, re, json
from typing import Optional
from openai import OpenAI
from env_util import load_env
load_env()
# ββ Provider registry βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
# Each entry: base_url, key_env, default_model, json_mode, extra_kwargs
PROVIDERS: dict[str, dict] = {
"github": {
"base_url": "https://models.github.ai/inference",
"key_env": "GITHUB_TOKEN",
"default_model": "openai/gpt-4o-mini",
"json_mode": True,
},
"openai": {
"base_url": None,
"key_env": "OPENAI_API_KEY",
"default_model": "gpt-4o-mini",
"json_mode": True,
},
"gemini": {
"base_url": "https://generativelanguage.googleapis.com/v1beta/openai/",
"key_env": "GEMINI_API_KEY",
"default_model": "gemini-2.0-flash",
"json_mode": True,
},
"mistral": {
"base_url": "https://api.mistral.ai/v1",
"key_env": "MISTRAL_API_KEY",
"default_model": "mistral-small-latest",
"json_mode": True,
},
"cohere": {
"base_url": "https://api.cohere.com/compatibility/v1",
"key_env": "COHERE_API_KEY",
"default_model": "command-r7b-12-2024",
"json_mode": False, # prompt-based JSON
},
"nvidia": {
"base_url": "https://integrate.api.nvidia.com/v1",
"key_env": "NVIDIA_API_KEY",
"default_model": "meta/llama-3.1-8b-instruct",
"json_mode": False,
},
"cerebras": {
"base_url": "https://api.cerebras.ai/v1",
"key_env": "CEREBRAS_API_KEY",
"default_model": "zai-glm-4.7",
"json_mode": False,
"reasoning": True, # needs higher token budget for chain-of-thought
},
"cloudflare": {
"base_url": None, # built dynamically using CLOUDFLARE_ID
"key_env": "CLOUDFLARE_API_TOKEN",
"default_model": "@cf/meta/llama-3.1-8b-instruct",
"json_mode": False,
},
"huggingface": {
"base_url": "https://router.huggingface.co/featherless-ai/v1",
"key_env": "HUGGINGFACE_API_TOKEN",
"default_model": "meta-llama/Llama-3.1-8B-Instruct",
"json_mode": False,
},
"sambanova": {
"base_url": "https://api.sambanova.ai/v1",
"key_env": "SAMBANOVA_API_KEY",
"default_model": "Meta-Llama-3.3-70B-Instruct",
"json_mode": False,
},
"openrouter": {
"base_url": "https://openrouter.ai/api/v1",
"key_env": "OPENROUTER_API_KEY",
"default_model": "nvidia/nemotron-3-ultra-550b-a55b:free",
"json_mode": False,
},
"deepseek": {
"base_url": "https://api.deepseek.com/v1",
"key_env": "DEEPSEEK_API_KEY",
"default_model": "deepseek-chat",
"json_mode": True,
},
}
JSON_ENFORCE_SUFFIX = (
"\n\nCRITICAL: Your response must be a single valid JSON object. "
"No markdown fences, no commentary, no preamble. Start with { and end with }."
)
def get_client(provider: str) -> tuple[OpenAI, dict]:
"""Return (OpenAI client, provider_config)."""
if provider not in PROVIDERS:
raise ValueError(f"Unknown provider '{provider}'. Choose from: {list(PROVIDERS)}")
cfg = PROVIDERS[provider]
key_env = cfg.get("key_env")
api_key = os.environ.get(key_env, "") if key_env else "local"
if key_env and not api_key:
raise SystemExit(f"Missing {key_env} in .env for provider '{provider}'")
base_url = cfg["base_url"]
if provider == "cloudflare":
account_id = os.environ.get("CLOUDFLARE_ID", "")
if not account_id:
raise SystemExit("Missing CLOUDFLARE_ID in .env")
base_url = f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/v1"
kwargs: dict = {"api_key": api_key}
if base_url:
kwargs["base_url"] = base_url
return OpenAI(**kwargs), cfg
def _extract_json(text: str) -> str:
"""Strip markdown fences and extract the first {...} block."""
text = text.strip()
# strip ```json ... ``` or ``` ... ```
text = re.sub(r"^```(?:json)?\s*", "", text, flags=re.MULTILINE)
text = re.sub(r"\s*```$", "", text, flags=re.MULTILINE)
text = text.strip()
# find first { ... } spanning the whole string
start = text.find("{")
if start == -1:
return text
depth, end = 0, -1
for i, ch in enumerate(text[start:], start):
if ch == "{":
depth += 1
elif ch == "}":
depth -= 1
if depth == 0:
end = i
break
return text[start:end + 1] if end != -1 else text[start:]
def chat_json(
provider: str,
system: str,
user: str,
model: Optional[str] = None,
temperature: float = 0.1,
max_tokens: int = 4000,
) -> str:
"""
Single chat completion that returns a JSON string.
For providers with native json_mode, passes response_format=json_object.
For others, appends a JSON-enforcement instruction and strips fences from
the response. Either way the caller gets a raw JSON string to parse.
"""
client, cfg = get_client(provider)
chosen_model = model or cfg["default_model"]
use_json_mode = cfg.get("json_mode", False)
# reasoning models need a much larger token budget
if cfg.get("reasoning") and max_tokens < 2048:
max_tokens = 2048
sys_msg = system if use_json_mode else system + JSON_ENFORCE_SUFFIX
kwargs: dict = {
"model": chosen_model,
"messages": [{"role": "system", "content": sys_msg},
{"role": "user", "content": user}],
"temperature": temperature,
"max_tokens": max_tokens,
}
if use_json_mode:
kwargs["response_format"] = {"type": "json_object"}
resp = client.chat.completions.create(**kwargs)
raw = resp.choices[0].message.content or ""
return _extract_json(raw) if not use_json_mode else raw.strip()
def list_providers() -> list[str]:
return list(PROVIDERS.keys())
def default_model(provider: str) -> str:
return PROVIDERS[provider]["default_model"]
# ββ Quick smoke test ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
if __name__ == "__main__":
import time, argparse
ap = argparse.ArgumentParser(description="Ping all providers or a specific one.")
ap.add_argument("--provider", default=None, help="test one provider only")
ap.add_argument("--prompt", default="Reply with valid JSON: {\"status\": \"ok\", \"msg\": \"hello\"}")
args = ap.parse_args()
targets = [args.provider] if args.provider else list(PROVIDERS.keys())
print(f"\n {'Provider':<14} {'Model':<42} {'ms':>6} {'Result'}")
print(" " + "-" * 90)
for p in targets:
cfg = PROVIDERS[p]
mod = cfg["default_model"]
try:
t0 = time.time()
raw = chat_json(p, "You are a helpful assistant.", args.prompt)
ms = int((time.time() - t0) * 1000)
# try to parse to confirm it's valid JSON
parsed = json.loads(raw)
print(f" [OK] {p:<12} {mod:<42} {ms:>6}ms {str(parsed)[:60]}")
except SystemExit as e:
print(f" [--] {p:<12} {mod:<42} {'missing key -- skipped'}")
except Exception as e:
print(f" [XX] {p:<12} {mod:<42} {str(e)[:70]}")
|