File size: 2,489 Bytes
196be8a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f329d5d
 
c6023f6
 
 
 
f329d5d
 
 
 
 
196be8a
 
 
 
 
 
 
 
 
 
f329d5d
196be8a
 
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
"""LLM client using the OpenAI-compatible Chat Completions API.

Provider is chosen with the LLM_PROVIDER env var (default: groq). Each provider
reads its own API key env var. The model can be overridden with LLM_MODEL.
"""
import os
from openai import OpenAI

PROVIDERS = {
    "groq": {
        "base_url": "https://api.groq.com/openai/v1",
        "key_env": "GROQ_API_KEY",
        "default_model": "llama-3.3-70b-versatile",
    },
    "openai": {
        "base_url": None,  # SDK default
        "key_env": "OPENAI_API_KEY",
        "default_model": "gpt-4o-mini",
    },
    "gemini": {
        "base_url": "https://generativelanguage.googleapis.com/v1beta/openai/",
        "key_env": "GEMINI_API_KEY",
        "default_model": "gemini-2.0-flash",
    },
}


def _config():
    provider = os.getenv("LLM_PROVIDER", "groq").lower()
    if provider not in PROVIDERS:
        raise ValueError(
            f"Unknown LLM_PROVIDER '{provider}'. Choose from: {', '.join(PROVIDERS)}."
        )
    cfg = PROVIDERS[provider]
    api_key = os.getenv(cfg["key_env"])
    model = os.getenv("LLM_MODEL", cfg["default_model"])
    return cfg, api_key, model


def is_configured() -> bool:
    """True if an API key for the selected provider is present."""
    try:
        _, api_key, _ = _config()
        return bool(api_key)
    except ValueError:
        return False


def _messages(system_prompt, history, user_prompt):
    msgs = [{"role": "system", "content": system_prompt}]
    for m in (history or []):
        role, content = m.get("role"), m.get("content")
        if role in ("user", "assistant") and isinstance(content, str) and content.strip():
            msgs.append({"role": role, "content": content})   # drop metadata/options/etc.
    msgs.append({"role": "user", "content": user_prompt})
    return msgs


def generate(system_prompt: str, user_prompt: str, history=None,
             temperature: float = 0.1, max_tokens: int = 700) -> str:
    cfg, api_key, model = _config()
    if not api_key:
        raise RuntimeError(f"Missing API key. Set the {cfg['key_env']} environment variable.")
    client = OpenAI(api_key=api_key, base_url=cfg["base_url"]) if cfg["base_url"] \
        else OpenAI(api_key=api_key)
    resp = client.chat.completions.create(
        model=model,
        temperature=temperature,
        max_tokens=max_tokens,
        messages=_messages(system_prompt, history, user_prompt),
    )
    return (resp.choices[0].message.content or "").strip()