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]}")