RICS / backend /llm /openai_client.py
StormShadow308's picture
Add demo documentation and Docker setup for v2 report generation system
aad7814
Raw
History Blame Contribute Delete
10.1 kB
"""Thin OpenAI client wrapper used across the v2 backend.
Exposes chat helpers over ``client.chat.completions.create`` with:
* Configurable concurrency gate (``MAX_CONCURRENT_LLM_CALLS``)
* Exponential backoff retries on HTTP 429 rate limits
* Async wrappers (``*_async``) that acquire an ``asyncio.Semaphore`` before
delegating to the sync implementation in a worker thread
When no API key is configured, ``is_available`` is False and callers fall back to
deterministic behaviour so the system runs (and tests pass) offline.
"""
from __future__ import annotations
import asyncio
import json
import logging
import random
import re
import threading
import time
from collections.abc import Callable
from contextlib import asynccontextmanager, contextmanager
from typing import Any, TypeVar
from backend.config import settings
logger = logging.getLogger(__name__)
T = TypeVar("T")
_client = None
_thread_llm_semaphore: threading.Semaphore | None = None
_async_llm_semaphore: asyncio.Semaphore | None = None
_semaphore_limit: int | None = None
_semaphore_lock = threading.Lock()
def is_available() -> bool:
"""Whether an OpenAI API key is configured."""
return bool((settings.openai_api_key or "").strip())
def _default_timeout() -> float:
return float(settings.openai_request_timeout_seconds)
def pipeline_timeout() -> float:
"""Strict timeout for generation-pipeline LLM / embedding calls."""
return float(settings.openai_pipeline_timeout_seconds)
def max_concurrent_llm_calls() -> int:
return max(1, int(settings.max_concurrent_llm_calls))
def _resolve_timeout(timeout: float | None) -> float:
return float(timeout if timeout is not None else pipeline_timeout())
def _reset_rate_limit_state() -> None:
global _thread_llm_semaphore, _async_llm_semaphore, _semaphore_limit
with _semaphore_lock:
_thread_llm_semaphore = None
_async_llm_semaphore = None
_semaphore_limit = None
def _ensure_semaphores() -> None:
global _thread_llm_semaphore, _async_llm_semaphore, _semaphore_limit
limit = max_concurrent_llm_calls()
with _semaphore_lock:
if _thread_llm_semaphore is None or _semaphore_limit != limit:
_thread_llm_semaphore = threading.Semaphore(limit)
_async_llm_semaphore = None
_semaphore_limit = limit
def _get_thread_semaphore() -> threading.Semaphore:
_ensure_semaphores()
assert _thread_llm_semaphore is not None
return _thread_llm_semaphore
def _get_async_semaphore() -> asyncio.Semaphore:
global _async_llm_semaphore
_ensure_semaphores()
if _async_llm_semaphore is None:
_async_llm_semaphore = asyncio.Semaphore(max_concurrent_llm_calls())
return _async_llm_semaphore
@contextmanager
def _sync_llm_slot():
sem = _get_thread_semaphore()
sem.acquire()
try:
yield
finally:
sem.release()
@asynccontextmanager
async def _async_llm_slot():
async with _get_async_semaphore():
yield
def _get_client():
global _client
if _client is None:
from openai import OpenAI
_client = OpenAI(
api_key=settings.openai_api_key,
timeout=_default_timeout(),
)
return _client
def reset_client() -> None:
"""Reset the cached client and concurrency gates (tests / config reloads)."""
global _client
_client = None
_reset_rate_limit_state()
def _is_rate_limit_error(exc: BaseException) -> bool:
status = getattr(exc, "status_code", None)
if status == 429:
return True
try:
from openai import APIStatusError, RateLimitError
except ImportError:
return False
return isinstance(exc, (RateLimitError,)) or (
isinstance(exc, APIStatusError) and exc.status_code == 429
)
def _retry_after_seconds(exc: BaseException) -> float | None:
response = getattr(exc, "response", None)
headers = getattr(response, "headers", None) if response is not None else None
if not headers:
return None
raw = headers.get("retry-after") or headers.get("Retry-After")
if raw is None:
return None
try:
return float(raw)
except (TypeError, ValueError):
return None
def _run_with_rate_limit_retry(
operation: Callable[[], T],
*,
gate_held: bool = False,
) -> T:
"""Execute *operation* under the concurrency gate with 429 backoff retries."""
max_retries = max(0, int(settings.openai_rate_limit_max_retries))
base_delay = float(settings.openai_rate_limit_backoff_base_seconds)
max_delay = float(settings.openai_rate_limit_backoff_max_seconds)
attempt = 0
while True:
try:
if gate_held:
return operation()
with _sync_llm_slot():
return operation()
except Exception as exc: # noqa: BLE001 — classify provider rate limits
if not _is_rate_limit_error(exc) or attempt >= max_retries:
raise
attempt += 1
retry_after = _retry_after_seconds(exc)
if retry_after is not None:
delay = min(retry_after, max_delay)
else:
delay = min(base_delay * (2 ** (attempt - 1)), max_delay)
delay *= 0.5 + random.random() * 0.5
logger.warning(
"OpenAI rate limit (429); retry %s/%s in %.2fs",
attempt,
max_retries,
delay,
)
time.sleep(delay)
async def _run_async_llm(
sync_fn: Callable[..., T],
/,
*args: Any,
**kwargs: Any,
) -> T:
"""Acquire the async concurrency gate, then run *sync_fn* in a worker thread."""
async with _async_llm_slot():
kwargs["_gate_held"] = True
return await asyncio.to_thread(sync_fn, *args, **kwargs)
def _parse_llm_json(content: str) -> dict:
"""Parse an LLM JSON payload, stripping optional markdown code fences."""
text = (content or "").strip()
if text.startswith("```"):
text = re.sub(r"^```(?:json)?\s*", "", text, flags=re.IGNORECASE)
text = re.sub(r"\s*```$", "", text)
text = text.strip()
try:
parsed = json.loads(text)
except json.JSONDecodeError:
logger.warning("LLM response is not valid JSON after fence stripping.")
return {}
return parsed if isinstance(parsed, dict) else {}
def chat_text(
messages: list[dict],
*,
model: str | None = None,
temperature: float = 0.0,
max_tokens: int | None = None,
timeout: float | None = None,
_gate_held: bool = False,
) -> str:
"""Return the assistant text for a chat completion."""
def _call() -> str:
resp = _get_client().chat.completions.create(
model=model or settings.mapping_model,
messages=messages,
temperature=temperature,
max_tokens=max_tokens,
timeout=_resolve_timeout(timeout),
)
return (resp.choices[0].message.content or "").strip()
return _run_with_rate_limit_retry(_call, gate_held=_gate_held)
async def chat_text_async(
messages: list[dict],
*,
model: str | None = None,
temperature: float = 0.0,
max_tokens: int | None = None,
timeout: float | None = None,
) -> str:
"""Async chat completion with concurrency gate + 429 retry."""
return await _run_async_llm(
chat_text,
messages,
model=model,
temperature=temperature,
max_tokens=max_tokens,
timeout=timeout,
)
def chat_vision_json(
messages: list[dict],
*,
model: str | None = None,
max_tokens: int | None = None,
timeout: float | None = None,
_gate_held: bool = False,
) -> dict:
"""Vision-capable chat completion returning parsed JSON."""
def _call() -> dict:
kwargs: dict[str, Any] = {
"model": model or settings.vision_model,
"messages": messages,
"temperature": 0.0,
"max_tokens": max_tokens or settings.vision_max_tokens,
"response_format": {"type": "json_object"},
"timeout": _resolve_timeout(timeout),
}
resp = _get_client().chat.completions.create(**kwargs)
content = (resp.choices[0].message.content or "").strip()
return _parse_llm_json(content)
return _run_with_rate_limit_retry(_call, gate_held=_gate_held)
async def chat_vision_json_async(
messages: list[dict],
*,
model: str | None = None,
max_tokens: int | None = None,
timeout: float | None = None,
) -> dict:
return await _run_async_llm(
chat_vision_json,
messages,
model=model,
max_tokens=max_tokens,
timeout=timeout,
)
def chat_json(
messages: list[dict[str, str]],
*,
model: str | None = None,
temperature: float = 0.0,
max_tokens: int | None = None,
timeout: float | None = None,
_gate_held: bool = False,
) -> dict:
"""Return the parsed JSON object for a chat completion in JSON mode."""
def _call() -> dict:
resp = _get_client().chat.completions.create(
model=model or settings.mapping_model,
messages=messages,
temperature=temperature,
max_tokens=max_tokens,
response_format={"type": "json_object"},
timeout=_resolve_timeout(timeout),
)
content = (resp.choices[0].message.content or "").strip()
return _parse_llm_json(content)
return _run_with_rate_limit_retry(_call, gate_held=_gate_held)
async def chat_json_async(
messages: list[dict[str, str]],
*,
model: str | None = None,
temperature: float = 0.0,
max_tokens: int | None = None,
timeout: float | None = None,
) -> dict:
"""Async JSON chat completion with concurrency gate + 429 retry."""
return await _run_async_llm(
chat_json,
messages,
model=model,
temperature=temperature,
max_tokens=max_tokens,
timeout=timeout,
)