amplegest / agent /llm.py
Viney's picture
feat: multi-provider LLM support, prominent chat, design pass, and new analytics
7880373
Raw
History Blame Contribute Delete
6.62 kB
"""agent/llm.py β€” provider-agnostic LLM factory.
The only module in the codebase that knows about specific LLM providers
(Anthropic Claude vs OpenAI). Every other agent module (graph, chat_agent,
translate) goes through `make_chat_model` / `build_system_message` and never
imports `ChatAnthropic` / `ChatOpenAI` directly.
Security note (public HF Space): a `RunConfig.api_key`, when set, is a value
the visitor pasted into the sidebar for THIS session only. It must never be
cached (st.cache_*), logged, or written to disk β€” see dashboard/model_picker.py
and app.py for how the config is threaded through as a per-run snapshot.
"""
from __future__ import annotations
import os
from dataclasses import dataclass
from typing import Optional
from langchain_core.messages import SystemMessage
ANTHROPIC_DEFAULT_MODEL = "claude-haiku-4-5-20251001"
OPENAI_DEFAULT_MODEL = "gpt-5-mini"
# (model_id, display_label) β€” data-driven so the sidebar picker and pricing
# table (agent/cost_log.py) stay in sync. Edit here to add/remove models.
MODEL_CATALOG: dict[str, list[tuple[str, str]]] = {
"anthropic": [
("claude-haiku-4-5-20251001", "Claude Haiku 4.5 β€” fast (default)"),
("claude-sonnet-4-6", "Claude Sonnet 4.6 β€” deeper analysis"),
("claude-opus-4-8", "Claude Opus 4.8 β€” highest quality"),
],
"openai": [
("gpt-5-mini", "GPT-5 mini β€” fast (default)"),
("gpt-5.1", "GPT-5.1 β€” deeper analysis"),
],
}
PROVIDER_LABELS: dict[str, str] = {
"anthropic": "Anthropic Claude",
"openai": "OpenAI",
}
_ENV_VARS: dict[str, str] = {
"anthropic": "ANTHROPIC_API_KEY",
"openai": "OPENAI_API_KEY",
}
_KEY_PREFIXES: dict[str, str] = {
"anthropic": "sk-ant-",
"openai": "sk-",
}
_CONSOLE_URLS: dict[str, str] = {
"anthropic": "console.anthropic.com",
"openai": "platform.openai.com",
}
@dataclass(frozen=True)
class RunConfig:
"""Frozen snapshot of provider/model/key for a single run.
Built once (in the main Streamlit thread, at the moment of the user
action) and passed as a plain argument into background threads and agent
functions β€” never re-read from st.session_state after construction, so a
mid-run change in the sidebar picker cannot affect an in-flight run.
"""
provider: str = "anthropic"
model: str = ANTHROPIC_DEFAULT_MODEL
api_key: Optional[str] = None # None β†’ the provider's own env var at construction
def default_config() -> RunConfig:
"""Today's behavior: Anthropic Haiku, key from ANTHROPIC_API_KEY env var."""
return RunConfig(provider="anthropic", model=ANTHROPIC_DEFAULT_MODEL, api_key=None)
def resolve_api_key(provider: str, user_key: Optional[str]) -> tuple[Optional[str], str]:
"""Resolve which key to use for *provider*.
Precedence: non-empty pasted key > env var > missing.
Returns (key_or_None, source) where source is one of "user", "env", "missing".
"""
user_key = (user_key or "").strip()
if user_key:
return user_key, "user"
env_key = os.environ.get(_ENV_VARS.get(provider, ""), "").strip()
if env_key:
return env_key, "env"
return None, "missing"
def key_looks_valid(provider: str, key: str) -> bool:
"""Cheap, no-network sanity check β€” catches pasting a key into the wrong provider."""
prefix = _KEY_PREFIXES.get(provider)
if not prefix or not key:
return True
return key.startswith(prefix)
def console_url(provider: str) -> str:
return _CONSOLE_URLS.get(provider, "")
def make_chat_model(
cfg: RunConfig,
*,
temperature: float = 0,
max_retries: int = 5,
max_tokens: Optional[int] = None,
):
"""Build a ChatAnthropic or ChatOpenAI client from *cfg*.
`api_key` is passed explicitly only when `cfg.api_key` is set β€” otherwise
the LangChain class reads its own provider env var, preserving today's
behavior exactly when no key was pasted in the UI.
"""
key, _source = resolve_api_key(cfg.provider, cfg.api_key)
kwargs: dict = {"model": cfg.model, "temperature": temperature, "max_retries": max_retries}
if max_tokens is not None:
kwargs["max_tokens"] = max_tokens
if key:
kwargs["api_key"] = key
if cfg.provider == "openai":
from langchain_openai import ChatOpenAI
return ChatOpenAI(**kwargs)
from langchain_anthropic import ChatAnthropic
return ChatAnthropic(**kwargs)
def build_system_message(
cfg: RunConfig,
primary_text: str,
extra_texts: Optional[list[str]] = None,
) -> SystemMessage:
"""Build a provider-appropriate SystemMessage.
Anthropic: content is a list of blocks; the primary block carries
`cache_control: {type: ephemeral}` so prompt caching keeps working.
Extra blocks (e.g. a language directive) are appended uncached, exactly
matching the byte layout the app used before this refactor.
OpenAI: plain string content β€” `cache_control` is an Anthropic-only
extension and would be ignored at best, rejected at worst.
"""
extra_texts = extra_texts or []
if cfg.provider == "openai":
parts = [primary_text] + list(extra_texts)
return SystemMessage(content="\n\n".join(parts))
content: list[dict] = [{
"type": "text",
"text": primary_text,
"cache_control": {"type": "ephemeral"},
}]
for text in extra_texts:
content.append({"type": "text", "text": text})
return SystemMessage(content=content)
_AUTH_MARKERS = ("authentication_error", "401", "invalid x-api-key", "incorrect api key", "invalid_api_key")
_OVERLOAD_MARKERS = ("overloaded_error", "529")
_RATE_LIMIT_MARKERS = ("rate_limit", "429", "quota")
def classify_llm_error(exc, provider: str) -> Optional[str]:
"""Map a raw exception/message to a friendly, provider-aware string.
Returns None when the error isn't recognized β€” the caller shows the raw
message in that case. Never interpolates the API key.
"""
text = str(exc).lower()
label = PROVIDER_LABELS.get(provider, provider)
if any(marker in text for marker in _AUTH_MARKERS):
return (
f"Your {label} API key was rejected (authentication error). "
f"Check it in Model settings β€” get a key at {console_url(provider)}."
)
if any(marker in text for marker in _OVERLOAD_MARKERS):
return f"{label} is temporarily overloaded. Please try again in a moment."
if any(marker in text for marker in _RATE_LIMIT_MARKERS):
return f"{label} rate limit reached. Please wait a moment and try again."
return None