Jordancole21's picture
Deploy native Gradio SEO page and tool
7e2c39a verified
Raw
History Blame Contribute Delete
7.14 kB
"""Shared OpenRouter inference adapter for all AI-powered Spaces.
- OpenAI-compatible chat-completions interface.
- Secrets only from environment (HF Space Secrets). Never logged.
- Structured JSON output validated against a spec (shared.validation);
one controlled repair attempt on invalid output.
- Bounded retries for transient failures only (429/5xx/timeouts),
honoring Retry-After. No automatic fallback on schema/safety failures;
deliberate fallback model only for provider-availability failures.
- Telemetry to stdout: provider, model, latency, usage, finish reason,
error class. Prompt content is never logged.
- LOREIFY_FAKE_AI=1 serves the caller-supplied fixture (local testing
without a key; never set this on a deployed Space).
"""
import json
import logging
import os
import time
import httpx
from .validation import validate
log = logging.getLogger("lf.ai")
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(name)s %(message)s")
class AIUnavailable(Exception):
"""User-safe failure. str(exc) is shown to the user."""
def _cfg():
return {
"api_key": os.environ.get("OPENROUTER_API_KEY", ""),
"base_url": os.environ.get("OPENROUTER_BASE_URL", "https://openrouter.ai/api/v1"),
"model": os.environ.get("OPENROUTER_MODEL", "google/gemini-3.6-flash"),
"fallback_model": os.environ.get("OPENROUTER_FALLBACK_MODEL", "openai/gpt-4o-mini"),
"timeout": float(os.environ.get("OPENROUTER_TIMEOUT_SECONDS", "60")),
# Generous default: reasoning-capable models (e.g. Gemini Flash) spend
# completion tokens on thinking before the JSON; a tight cap truncates.
"max_tokens": int(os.environ.get("OPENROUTER_MAX_TOKENS", "6000")),
"temperature": float(os.environ.get("OPENROUTER_TEMPERATURE", "0.9")),
}
USER_ERROR = (
"The generator hit a snag and couldn't finish. Nothing was charged to you — "
"please try again in a moment."
)
NO_KEY_ERROR = (
"The generator isn't configured yet (missing API credentials). "
"If you run this Space, add the OPENROUTER_API_KEY secret."
)
def _post_completion(model: str, messages: list, cfg: dict, max_tokens: int):
"""One provider call with bounded retries for transient failures."""
headers = {
"Authorization": f"Bearer {cfg['api_key']}",
"HTTP-Referer": "https://loreify.ai",
"X-Title": "Loreify D&D Tools",
"Content-Type": "application/json",
}
body = {
"model": model,
"messages": messages,
"max_tokens": max_tokens,
"temperature": cfg["temperature"],
"response_format": {"type": "json_object"},
}
delay = 1.5
last_exc = None
for attempt in range(3):
start = time.time()
try:
resp = httpx.post(
f"{cfg['base_url']}/chat/completions",
headers=headers, json=body, timeout=cfg["timeout"],
)
except (httpx.TimeoutException, httpx.TransportError) as exc:
last_exc = exc
log.info("call model=%s attempt=%d error_class=%s", model, attempt, type(exc).__name__)
time.sleep(delay)
delay *= 2
continue
latency = time.time() - start
if resp.status_code == 429 or resp.status_code >= 500:
last_exc = AIUnavailable(USER_ERROR)
retry_after = resp.headers.get("retry-after")
wait = min(float(retry_after), 15.0) if retry_after and retry_after.isdigit() else delay
log.info("call model=%s attempt=%d status=%d retry_in=%.1f", model, attempt, resp.status_code, wait)
time.sleep(wait)
delay *= 2
continue
if resp.status_code != 200:
log.info("call model=%s status=%d error_class=http", model, resp.status_code)
raise AIUnavailable(USER_ERROR)
data = resp.json()
choice = (data.get("choices") or [{}])[0]
usage = data.get("usage") or {}
log.info(
"call ok provider=openrouter model=%s latency=%.2fs prompt_tokens=%s "
"completion_tokens=%s finish=%s",
data.get("model", model), latency,
usage.get("prompt_tokens"), usage.get("completion_tokens"),
choice.get("finish_reason"),
)
content = (choice.get("message") or {}).get("content") or ""
if not content:
raise AIUnavailable(USER_ERROR)
return content
raise last_exc if isinstance(last_exc, AIUnavailable) else AIUnavailable(USER_ERROR)
def _parse_json(text: str):
text = text.strip()
if text.startswith("```"):
text = text.strip("`")
if text.startswith("json"):
text = text[4:]
start, end = text.find("{"), text.rfind("}")
if start == -1 or end == -1:
raise ValueError("no JSON object found")
return json.loads(text[start : end + 1])
def generate_json(
system_prompt: str,
user_prompt: str,
spec: dict,
fixture: dict | None = None,
max_tokens: int | None = None,
) -> dict:
"""Generate a spec-valid JSON object. Raises AIUnavailable on failure."""
if os.environ.get("LOREIFY_FAKE_AI") == "1" and fixture is not None:
time.sleep(0.4) # simulate latency so loading states are testable
return fixture
cfg = _cfg()
if not cfg["api_key"]:
raise AIUnavailable(NO_KEY_ERROR)
tokens = min(max_tokens or cfg["max_tokens"], cfg["max_tokens"])
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt},
]
models = [cfg["model"]] + ([cfg["fallback_model"]] if cfg["fallback_model"] else [])
content = None
for i, model in enumerate(models):
try:
content = _post_completion(model, messages, cfg, tokens)
break
except AIUnavailable:
if i == len(models) - 1:
raise
log.info("falling back model=%s -> %s", model, models[i + 1])
active_model = model
# Parse + validate, with one controlled repair attempt.
for attempt in ("primary", "repair"):
try:
data = _parse_json(content)
problems = validate(data, spec)
except (ValueError, json.JSONDecodeError) as exc:
data, problems = None, [f"invalid JSON: {exc}"]
if not problems:
return data
log.info("schema attempt=%s problems=%d", attempt, len(problems))
if attempt == "repair":
raise AIUnavailable(USER_ERROR)
repair_messages = messages + [
{"role": "assistant", "content": content},
{
"role": "user",
"content": (
"Your previous response did not match the required JSON schema. "
"Problems: " + "; ".join(problems[:10]) + ". "
"Respond again with ONLY the corrected JSON object, no prose."
),
},
]
content = _post_completion(active_model, repair_messages, cfg, tokens)
raise AIUnavailable(USER_ERROR) # unreachable