Terminal / backend /models /ai_client.py
Pulka
fix: disable HF provider by default (402 credits), 2.5s backoff on 429
2dc0681 verified
Raw
History Blame Contribute Delete
14.8 kB
"""
ai_client.py — Cloud AI Client for free remote deployments
Backend-first model router for mobile/no-PC usage. It prefers free or low-cost
OpenAI-compatible providers when configured and falls back across providers before
returning an error. The iPhone remains only a control surface; all model calls run
from the deployed backend/HF Space.
v2 — timeout + asyncio.to_thread per ogni chiamata sync (non blocca l'event loop),
fallback immediato al provider successivo su timeout o errore.
"""
from __future__ import annotations
import asyncio
import os
from dataclasses import dataclass
from typing import AsyncIterator, Optional
from openai import OpenAI
# Timeout per provider — abbassabile via env per reti lente
PROVIDER_TIMEOUT: float = float(os.getenv("PROVIDER_TIMEOUT", "15"))
STREAM_TIMEOUT: float = float(os.getenv("STREAM_TIMEOUT", "30"))
@dataclass(frozen=True)
class ProviderConfig:
name: str
api_key: str
base_url: str
default_model: str
embedding_model: Optional[str] = None
class AIClient:
def __init__(self) -> None:
self.providers = self._discover_providers()
self.provider_name = self.providers[0].name if self.providers else "unconfigured"
self.default_model = self.providers[0].default_model if self.providers else os.getenv("LLM_MODEL", "openrouter/auto")
self.client = self._client_for(self.providers[0]) if self.providers else None
# ── Discovery ────────────────────────────────────────────────────────────
def _discover_providers(self) -> list[ProviderConfig]:
"""
S192 FIX: ordine provider ottimizzato per HF Space.
Groq FIRST (llama-3.1-8b-instant — gratuito, veloce, GROQ_API_KEY presente in Space).
Gemini second (gemini-2.0-flash — key GEMINI_API_KEY).
OpenRouter third. HF LAST e solo se DISABLE_HF_PROVIDER non è settato.
"""
providers: list[ProviderConfig] = []
# 1. Groq first — key presente in HF Space, llama-3.1-8b-instant < 5s
groq_key = os.getenv("GROQ_API_KEY")
if groq_key:
providers.append(ProviderConfig(
name="groq",
api_key=groq_key,
base_url="https://api.groq.com/openai/v1",
# S192: usa 8b-instant (stesso di smolagents) — non 70b-versatile (premium)
default_model=os.getenv("GROQ_MODEL", os.getenv("LLM_MODEL_GROQ", "llama-3.1-8b-instant")),
))
# 2. Gemini second
gemini_key = os.getenv("GEMINI_API_KEY") or os.getenv("GOOGLE_API_KEY")
if gemini_key:
providers.append(ProviderConfig(
name="gemini",
api_key=gemini_key,
base_url="https://generativelanguage.googleapis.com/v1beta/openai/",
default_model=os.getenv("GEMINI_MODEL", os.getenv("LLM_MODEL_GEMINI", "gemini-2.0-flash")),
))
# 3. OpenRouter third
openrouter_key = os.getenv("OPENROUTER_API_KEY")
if openrouter_key:
providers.append(ProviderConfig(
name="openrouter",
api_key=openrouter_key,
base_url="https://openrouter.ai/api/v1",
default_model=os.getenv("OPENROUTER_MODEL", os.getenv("LLM_MODEL", "meta-llama/llama-3.3-70b-instruct:free")),
))
# 4. HF router LAST — disabilitato di default (crediti 402 esauriti sul free tier)
# Abilita solo se ENABLE_HF_PROVIDER=1 è settato esplicitamente
hf_key = os.getenv("HF_TOKEN") or os.getenv("HUGGINGFACE_API_KEY") or os.getenv("HUGGINGFACE_TOKEN")
if hf_key and os.getenv("ENABLE_HF_PROVIDER"):
providers.append(ProviderConfig(
name="huggingface",
api_key=hf_key,
base_url=os.getenv("HF_OPENAI_BASE_URL", "https://router.huggingface.co/v1"),
default_model=os.getenv("HF_MODEL", os.getenv("LLM_MODEL_HF", "Qwen/Qwen2.5-Coder-32B-Instruct")),
))
# 5. OpenAI compat last
openai_key = os.getenv("OPENAI_API_KEY")
if openai_key:
providers.append(ProviderConfig(
name="openai_compatible",
api_key=openai_key,
base_url=os.getenv("OPENAI_API_BASE", "https://api.openai.com/v1"),
default_model=os.getenv("OPENAI_MODEL", os.getenv("LLM_MODEL", "gpt-4o-mini")),
embedding_model=os.getenv("EMBEDDING_MODEL", "text-embedding-3-small"),
))
return providers
# ── Client factory ────────────────────────────────────────────────────────
def _client_for(self, provider: ProviderConfig) -> OpenAI:
"""Istanzia il client con timeout esplicito — mai blocca l'event loop."""
return OpenAI(
api_key=provider.api_key,
base_url=provider.base_url,
timeout=PROVIDER_TIMEOUT,
max_retries=0, # gestiamo noi il fallback tra provider
)
def _model_for(self, provider: ProviderConfig, requested: Optional[str]) -> str:
return requested or provider.default_model
# ── Health ────────────────────────────────────────────────────────────────
async def health(self) -> dict:
if not self.providers:
return {
"available": False,
"provider": "none",
"error": "No remote provider key configured. Set OPENROUTER_API_KEY, GEMINI_API_KEY, GROQ_API_KEY, HF_TOKEN or OPENAI_API_KEY.",
"models": [],
}
checks: list[dict] = []
for provider in self.providers:
client = self._client_for(provider)
try:
await asyncio.wait_for(
asyncio.to_thread(
client.chat.completions.create,
model=provider.default_model,
messages=[{"role": "user", "content": "ping"}],
max_tokens=1,
temperature=0,
),
timeout=PROVIDER_TIMEOUT,
)
checks.append({"provider": provider.name, "available": True, "model": provider.default_model})
self.provider_name = provider.name
self.default_model = provider.default_model
self.client = client
return {
"available": True,
"provider": provider.name,
"models": [p.default_model for p in self.providers],
"default": provider.default_model,
"checks": checks,
}
except Exception as exc:
checks.append({"provider": provider.name, "available": False, "model": provider.default_model, "error": str(exc)})
return {
"available": False,
"provider": "configured_but_unavailable",
"models": [p.default_model for p in self.providers],
"checks": checks,
}
# ── Chat (non-stream) ─────────────────────────────────────────────────────
async def chat(
self,
messages: list,
*,
model: Optional[str] = None,
temperature: float = 0.7,
max_tokens: int = 4096,
timeout: Optional[float] = None,
) -> str:
"""
Itera sui provider configurati in ordine di priorità.
Passa al successivo se il provider corrente supera il timeout o restituisce errore.
Ogni chiamata sync viene eseguita in asyncio.to_thread — mai blocca l'event loop.
"""
per_provider_timeout = timeout or PROVIDER_TIMEOUT
last_error: Exception | None = None
for provider in self.providers:
client = self._client_for(provider)
# Retry up to 2 times per provider (handles transient 429 burst)
for attempt in range(2):
try:
response = await asyncio.wait_for(
asyncio.to_thread(
client.chat.completions.create,
model=self._model_for(provider, model),
messages=messages,
temperature=temperature,
max_tokens=max_tokens,
stream=False,
),
timeout=per_provider_timeout,
)
# Aggiorna stato attivo
self.provider_name = provider.name
self.default_model = self._model_for(provider, model)
self.client = client
return response.choices[0].message.content or ""
except asyncio.TimeoutError:
last_error = TimeoutError(f"{provider.name} non ha risposto entro {per_provider_timeout}s")
break # timeout → prossimo provider senza retry
except Exception as exc:
last_error = exc
exc_str = str(exc)
is_rate_limit = "429" in exc_str or "rate_limit" in exc_str.lower()
is_no_credits = "402" in exc_str or "depleted" in exc_str.lower()
if is_no_credits:
break # crediti esauriti → skip immediato
if is_rate_limit and attempt == 0:
# Primo 429: aspetta 2.5s e riprova (window TPM si svuota)
await asyncio.sleep(2.5)
continue
break # secondo tentativo o altro errore → prossimo provider
raise RuntimeError(f"Nessun provider disponibile: {last_error}")
# ── Stream chat ───────────────────────────────────────────────────────────
async def stream_chat(
self,
messages: list,
*,
model: Optional[str] = None,
temperature: float = 0.7,
max_tokens: int = 4096,
) -> AsyncIterator[str]:
"""
Streaming con fallback automatico tra provider.
FIX(audit): l'iterazione dei chunk OpenAI è sincrona (blocking HTTP read).
Raccogliamo tutti i chunk in asyncio.to_thread per non bloccare il
FastAPI event loop. L'output rimane un async iterator verso il caller.
"""
last_error: Exception | None = None
for provider in self.providers:
client = self._client_for(provider)
try:
# Streaming reale: thread producer → asyncio.Queue → yield real-time
# BUG-8 fix: prima raccoglievamo TUTTO prima di yieldare (falso streaming)
q: asyncio.Queue[str | None] = asyncio.Queue()
_loop = asyncio.get_running_loop() # S166-Fix1: cattura nel contesto async
def _stream_to_queue() -> None:
try:
stream = client.chat.completions.create(
model=self._model_for(provider, model),
messages=messages,
temperature=temperature,
max_tokens=max_tokens,
stream=True,
)
for chunk in stream:
if chunk.choices and chunk.choices[0].delta.content:
_loop.call_soon_threadsafe( # S166-Fix1
q.put_nowait, chunk.choices[0].delta.content
)
finally:
_loop.call_soon_threadsafe(q.put_nowait, None) # S166-Fix1
import threading
t = threading.Thread(target=_stream_to_queue, daemon=True)
t.start()
self.provider_name = provider.name
self.default_model = self._model_for(provider, model)
self.client = client
deadline = _loop.time() + STREAM_TIMEOUT # S166-Fix1
while True:
remaining = deadline - _loop.time() # S166-Fix1
if remaining <= 0:
raise TimeoutError(f"stream timeout {STREAM_TIMEOUT}s")
try:
token = await asyncio.wait_for(q.get(), timeout=min(remaining, 5.0))
except asyncio.TimeoutError:
raise TimeoutError(f"stream timeout {STREAM_TIMEOUT}s")
if token is None:
break
yield token
return
except asyncio.TimeoutError:
last_error = TimeoutError(f"{provider.name} stream timeout dopo {STREAM_TIMEOUT}s")
continue
except Exception as exc:
last_error = exc
exc_str = str(exc)
if "429" in exc_str or "402" in exc_str or "rate_limit" in exc_str.lower():
continue
continue
raise RuntimeError(f"Nessun provider streaming disponibile: {last_error}")
# ── Embeddings ────────────────────────────────────────────────────────────
async def embed(self, text: str, model: str = "text-embedding-3-small") -> list[float]:
for provider in self.providers:
embedding_model = provider.embedding_model or os.getenv("EMBEDDING_MODEL")
if not embedding_model:
continue
client = self._client_for(provider)
try:
response = await asyncio.wait_for(
asyncio.to_thread(
client.embeddings.create,
input=[text],
model=embedding_model or model,
),
timeout=PROVIDER_TIMEOUT,
)
return response.data[0].embedding
except Exception:
continue
return []