Terminal / models /ai_client.py
Baida07's picture
sync: 166 file da Baida98/AI@a6ac2424e11e5c320c5ff688e1ce7addac64cdab (local-fallback deploy-all) (#47)
ed28aa2
Raw
History Blame
25.2 kB
"""
ai_client.py — Real Provider Fleet (LLM API Ensemble)
Legge la flotta di provider LLM da Supabase (tabella `ai_providers`, se disponibile)
con fallback su variabili d'ambiente per gli 8 provider realmente configurati nel
progetto (Groq, OpenRouter, Cerebras, SambaNova, Gemini, NVIDIA NIM, OpenAI, HF Router).
STORIA: la versione precedente ("Dynamic SQL Fleet Edition") assumeva una flotta di
10 ZeroGPU Spaces (HF_SPACE_A1..E2) mai realmente esistita — nessuna di quelle env var
è mai stata configurata né documentata altrove nel repo, quindi self.providers era
sempre vuoto in produzione (ogni chiamata falliva silenziosamente con
"tutti i nodi A-E sono offline"). Sostituito con i provider realmente attivi
(vedi ENV_AUDIT_P21.md §6 e artifacts/agente-ai/src/config/env.ts).
"""
from __future__ import annotations
import asyncio
import json
import os
import re
import time as _time_mod
from dataclasses import dataclass
from typing import AsyncIterator, Optional, List, Tuple
from openai import OpenAI
from api.semantic_cache import get_cached_response, set_cached_response
import logging
_logger = logging.getLogger("agente_ai")
class ProviderUnavailableError(RuntimeError):
"""Raised when no configured LLM provider can produce a response.
The error deliberately includes provider names only, never credentials or
raw upstream payloads, so callers can distinguish infrastructure failure
from a model answer without leaking sensitive data.
"""
def __init__(self, providers: list[str] | tuple[str, ...]) -> None:
self.providers = tuple(providers)
detail = ", ".join(self.providers) if self.providers else "none"
super().__init__(f"provider_unavailable: {detail}")
@dataclass(frozen=True)
class ProviderConfig:
id: int = 0
name: str = ""
api_key: str = ""
base_url: str = ""
default_model: str = ""
tier: int = 1
purpose: str = "reasoning"
profile: str = "general"
@property
def identity(self) -> tuple[str, str, str]:
"""Stable identity: different profiles must never share a client cache entry."""
return (self.name, self.profile, self.base_url)
# Definizione statica dei provider LLM realmente attivi nel progetto.
# base_url punta sempre all'endpoint OpenAI-compatible ufficiale del provider
# (nessun proxy CF Worker qui: questo client gira lato backend Python, non browser).
_PROVIDER_DEFS = [
# tier 0 — free tier veloce e affidabile
{"name": "groq", "env_key": "GROQ_API_KEY", "base_url": "https://api.groq.com/openai/v1", "model_env": "GROQ_MODEL", "default_model": "openai/gpt-oss-120b", "tier": 0, "purpose": "reasoning"},
{"name": "cerebras", "env_key": "CEREBRAS_API_KEY", "base_url": "https://api.cerebras.ai/v1", "model_env": "CEREBRAS_MODEL", "default_model": "gpt-oss-120b", "tier": 0, "purpose": "reasoning"},
{"name": "sambanova", "env_key": "SAMBANOVA_API_KEY", "base_url": "https://api.sambanova.ai/v1", "model_env": "SAMBANOVA_MODEL", "default_model": "DeepSeek-V3.1", "tier": 0, "purpose": "reasoning"},
# tier 1 — free tier con rate limit più stretti
{"name": "openrouter", "env_key": "OPENROUTER_API_KEY", "base_url": "https://openrouter.ai/api/v1", "model_env": "OPENROUTER_MODEL","default_model": "openai/gpt-oss-20b:free", "tier": 1, "purpose": "coding"},
{"name": "hf_router", "env_key": "HF_TOKEN", "base_url": "https://router.huggingface.co/v1", "model_env": "HF_MODEL", "default_model": "Qwen/Qwen2.5-Coder-32B-Instruct", "tier": 1, "purpose": "coding"},
{"name": "gemini", "env_key": "GEMINI_API_KEY", "base_url": "https://generativelanguage.googleapis.com/v1beta/openai/", "model_env": "GEMINI_MODEL", "default_model": "gemini-3.6-flash", "tier": 1, "purpose": "memory"},
# tier 2 — fallback opzionale (spesso a pagamento o quota limitata)
{"name": "nvidia", "env_key": "NVIDIA_API_KEY", "base_url": "https://integrate.api.nvidia.com/v1", "model_env": "NVIDIA_MODEL", "default_model": "nvidia/nemotron-3-ultra-550b-a55b", "tier": 2, "purpose": "audit"},
]
class AIClient:
def __init__(self) -> None:
self.providers = self._load_providers()
self._client_cache: dict[tuple[str, str, str], OpenAI] = {}
# Round-robin e circuit breaker sono indicizzati per purpose e profilo.
self._rr_indices: dict[str, int] = {}
self._breaker: dict[tuple[str, str, str], dict[str, float | int]] = {}
self._breaker_threshold = 2
self._breaker_cooldown_s = 60.0
def _load_providers(self) -> list[ProviderConfig]:
"""Carica la flotta: prova Supabase (tabella `ai_providers`, source of
truth dichiarata in supabase/migrations/20260711_ai_providers_fleet.sql),
fallback sui provider reali via env se Supabase non è raggiungibile/vuoto
(es. progetto sospeso per fatturazione, tabella non ancora popolata)."""
database_providers = self._try_load_from_supabase()
environment_profiles = [
profile
for definition in _PROVIDER_DEFS
for profile in self._profile_rows_from_env(definition)
]
if environment_profiles:
profiled_names = {profile.name for profile in environment_profiles}
# Explicit profile pools override a same-provider Supabase credential;
# DB providers not covered by a pool remain available as fallbacks.
database_providers = [
provider for provider in database_providers
if provider.name not in profiled_names
]
return environment_profiles + database_providers
return database_providers or self._discover_providers_from_env()
@staticmethod
def _runtime_model_override(row: dict) -> str:
"""Use an explicit model environment override for a known provider endpoint.
Supabase remains the source for provider credentials and ordering; runtime
model-selection variables deliberately win so emergency model migrations
do not require reading or mutating provider secrets in the database.
"""
database_model = str(row.get("default_model", ""))
row_base_url = str(row.get("base_url", "")).rstrip("/")
for definition in _PROVIDER_DEFS:
if row_base_url == definition["base_url"].rstrip("/"):
return os.getenv(definition["model_env"], database_model)
return database_model
@staticmethod
def _is_legacy_schema_error(exc: Exception) -> bool:
"""Riconosce il layout `ai_providers` precedente alla flotta canonica.
Quel layout espone `model_name`, `priority` e `provider_type`, ma
contiene record storici e modelli deprecati. Fino alla migrazione non va
promosso a source of truth: il fallback ambiente aggiornato è più sicuro.
"""
message = str(exc).lower()
return (
"column ai_providers." in message
and "does not exist" in message
and any(column in message for column in (
"default_model", "tier", "purpose", "success_count",
))
)
@staticmethod
def _profile_rows_from_env(definition: dict) -> list[ProviderConfig]:
"""Load optional per-provider profiles without logging secret values.
Format: ``<PROVIDER>_PROFILES_JSON=[{"profile":"p1","api_key":"...", "model":"..."}]``.
The legacy single-key variable remains supported and is loaded after profiles.
"""
env_name = f"{definition['name'].upper()}_PROFILES_JSON"
raw = os.getenv(env_name, "").strip()
if not raw:
return []
try:
rows = json.loads(raw)
except json.JSONDecodeError:
_logger.warning("AIClient: %s non valido, profili ignorati", env_name)
return []
if not isinstance(rows, list):
_logger.warning("AIClient: %s deve essere un array JSON", env_name)
return []
result: list[ProviderConfig] = []
for index, row in enumerate(rows):
if not isinstance(row, dict) or not row.get("api_key"):
continue
result.append(ProviderConfig(
id=-(index + 1),
name=definition["name"],
api_key=str(row["api_key"]),
base_url=str(row.get("base_url") or definition["base_url"]),
default_model=str(row.get("model") or os.getenv(definition["model_env"], definition["default_model"])),
tier=definition["tier"],
purpose=str(row.get("purpose") or definition["purpose"]),
profile=str(row.get("profile") or f"profile-{index + 1}"),
))
return result
def _try_load_from_supabase(self) -> list[ProviderConfig]:
url = os.getenv("SUPABASE_URL", "")
key = os.getenv("SUPABASE_SERVICE_ROLE_KEY", "")
if not url or not key:
return []
try:
from supabase import create_client
sb = create_client(url, key)
try:
res = (
sb.table("ai_providers")
.select("id,name,api_key,base_url,default_model,tier,purpose")
.eq("is_active", True)
.order("tier", desc=False)
.order("success_count", desc=True)
.execute()
)
except Exception as exc:
if self._is_legacy_schema_error(exc):
# Non usare il layout storico: contiene provider fittizi e
# modelli superati. La migrazione normalizzerà la tabella;
# nel frattempo il caller seleziona il fallback env corrente.
return []
raise
rows = res.data or []
return [
ProviderConfig(
id=row["id"], name=row["name"], api_key=row["api_key"],
base_url=row["base_url"], default_model=self._runtime_model_override(row),
tier=row["tier"], purpose=row["purpose"],
# Legacy schema has no profile column: the row id is still a
# stable profile identity and prevents client-cache collisions.
profile=f"db-{row['id']}",
)
for row in rows
]
except Exception as e:
_logger.warning(f"AIClient: Supabase ai_providers non disponibile, uso fallback env: {e}")
return []
def _discover_providers_from_env(self) -> list[ProviderConfig]:
"""Discovery dei provider LLM realmente configurati via env var (vedi
ENV_AUDIT_P21.md §6). Un provider viene incluso solo se la sua API key
è impostata — nessun placeholder, nessun nodo fantasma."""
providers = []
for i, d in enumerate(_PROVIDER_DEFS):
providers.extend(self._profile_rows_from_env(d))
api_key = os.getenv(d["env_key"], "")
if not api_key:
continue
providers.append(ProviderConfig(
id=i,
name=d["name"],
api_key=api_key,
base_url=d["base_url"],
default_model=os.getenv(d["model_env"], d["default_model"]),
tier=d["tier"],
purpose=d["purpose"],
profile="legacy",
))
if not providers:
_logger.error("AIClient: nessuna API key provider configurata (Groq/OpenRouter/Cerebras/SambaNova/Gemini/NVIDIA/HF_TOKEN tutte assenti)")
return providers
def _client_for(self, provider: ProviderConfig) -> OpenAI:
if provider.identity not in self._client_cache:
self._client_cache[provider.identity] = OpenAI(
api_key=provider.api_key,
base_url=provider.base_url,
# I task coding possono richiedere più di 20 s prima del primo
# chunk dal fallback gratuito; il budget esterno resta finito.
timeout=45,
max_retries=0
)
return self._client_cache[provider.identity]
def _is_available(self, provider: ProviderConfig) -> bool:
state = self._breaker.get(provider.identity)
return not state or float(state.get("open_until", 0.0)) <= _time_mod.monotonic()
def _record_success(self, provider: ProviderConfig) -> None:
self._breaker.pop(provider.identity, None)
def _record_failure(self, provider: ProviderConfig, exc: Exception) -> None:
message = str(exc).lower()
if not any(token in message for token in ("401", "403", "429", "500", "502", "503", "504", "rate limit", "quota")):
return
state = self._breaker.setdefault(provider.identity, {"failures": 0, "open_until": 0.0})
failures = int(state.get("failures", 0)) + 1
severe = any(token in message for token in ("401", "403"))
quota_limited = any(token in message for token in ("429", "rate limit", "quota"))
# A quota/rate-limit response is deterministic: retrying the same
# profile immediately only creates a storm. Open that profile on the
# first signal and let the provider pool move to another provider.
threshold = 1 if severe or quota_limited else self._breaker_threshold
if failures >= threshold:
cooldown = 900.0 if severe else self._rate_limit_cooldown_seconds(message) if quota_limited else self._breaker_cooldown_s
state["open_until"] = _time_mod.monotonic() + cooldown
state["failures"] = failures
@staticmethod
def _rate_limit_cooldown_seconds(message: str) -> float:
"""Return a provider reset-aware cooldown, never shorter than 15 min."""
reset_match = re.search(r"x-ratelimit-reset[^0-9]*(\d{10,13})", message, re.IGNORECASE)
if reset_match:
reset_value = float(reset_match.group(1))
reset_epoch = reset_value / 1000.0 if reset_value > 10_000_000_000 else reset_value
return max(900.0, reset_epoch - _time_mod.time())
return 900.0
def _execution_pool(self, providers: list[ProviderConfig], purpose: str) -> list[ProviderConfig]:
"""Return one rotated, healthy profile per provider endpoint group."""
groups: dict[tuple[str, str], list[ProviderConfig]] = {}
for provider in providers:
if not self._is_available(provider):
continue
groups.setdefault((provider.name, provider.base_url), []).append(provider)
selected: list[ProviderConfig] = []
for group_key, profiles in groups.items():
index_key = f"{purpose}:{group_key[0]}:{group_key[1]}"
start = self._rr_indices.get(index_key, 0)
selected.append(profiles[start % len(profiles)])
self._rr_indices[index_key] = start + 1
return selected
def _inter_provider_fallback_pool(
self,
purpose: str,
excluded: set[str] | None = None,
providers: list[ProviderConfig] | None = None,
) -> list[ProviderConfig]:
"""Select one healthy profile per provider, prioritizing the target purpose.
A provider whose complete profile group is open in the circuit breaker is
absent from this list; the next healthy provider becomes the automatic
fallback. This prevents retry storms against an exhausted pool.
"""
excluded = excluded or set()
source = self.providers if providers is None else providers
candidates = [
provider for provider in source
if provider.name not in excluded and self._is_available(provider)
]
candidates.sort(key=lambda provider: (
0 if provider.purpose == purpose else 1,
provider.tier,
provider.name,
provider.profile,
))
return self._execution_pool(candidates, f"fallback:{purpose}")
async def _fetch_one(self, provider: ProviderConfig, messages: list, temperature: float, max_tokens: int) -> Tuple[ProviderConfig, str, float]:
start = _time_mod.monotonic()
try:
client = self._client_for(provider)
response = await asyncio.wait_for(
asyncio.to_thread(
client.chat.completions.create,
model=provider.default_model,
messages=messages,
temperature=temperature,
max_tokens=max_tokens
),
# Il fallback non-streaming deve avere lo stesso budget del client:
# 15s scartava provider sani su richieste coding che richiedono
# più tempo per produrre una risposta completa dopo uno stream interrotto.
timeout=45
)
self._record_success(provider)
return provider, response.choices[0].message.content or "", _time_mod.monotonic() - start
except Exception as e:
self._record_failure(provider, e)
_logger.warning(f"Provider {provider.name}/{provider.profile} fallito: {e}")
return provider, f"ERROR: {str(e)}", 0.0
def _get_round_robin_provider(self, purpose: str) -> Optional[ProviderConfig]:
"""Ritorna uno dei provider con lo stesso purpose usando round-robin."""
purpose_providers = [p for p in self.providers if p.purpose == purpose]
if not purpose_providers:
return None
idx = self._rr_indices.get(purpose, 0)
provider = purpose_providers[idx % len(purpose_providers)]
self._rr_indices[purpose] = idx + 1
return provider
async def chat(self, messages: list, model: Optional[str] = None, temperature: float = 0.7, max_tokens: int = 4096) -> str:
if not self.providers: self.providers = self._load_providers()
# S-CACHE-1: Lookup semantico preventivo
cached = await get_cached_response(messages)
if cached:
_logger.info("[fleet] Semantic Cache HIT — risparmio quota provider")
return cached
# 1. Determina l'intento per il routing iniziale (specializzazione per purpose,
# non per nodo hardware: i provider reali non sono dedicati per profilo).
prompt = str(messages[-1].get("content", "")).lower()
if any(k in prompt for k in ["code", "python", "fix", "script"]):
primary_purpose = "coding"
elif any(k in prompt for k in ["cerca", "ricorda", "database", "memory"]):
primary_purpose = "memory"
elif any(k in prompt for k in ["sicuro", "audit", "verifica", "check"]):
primary_purpose = "audit"
else:
primary_purpose = "reasoning"
# 2. Costruisce il pool di esecuzione (provider del purpose target + fallback globale)
pool = [p for p in self.providers if p.purpose == primary_purpose]
# Se il pool è vuoto, prendiamo i primi 4 provider disponibili (per tier)
if not pool:
pool = self.providers[:4]
if not pool:
raise ProviderUnavailableError([])
# 3. Un solo profilo per endpoint e richiesta: round-robin evita che
# profili condividano quota e client, mentre provider diversi restano
# disponibili come ensemble/fallback.
pool = self._execution_pool(pool, primary_purpose)
results = []
if pool:
tasks = [self._fetch_one(p, messages, temperature, max_tokens) for p in pool]
results = await asyncio.gather(*tasks)
valid = [result for result in results if not result[1].startswith("ERROR:") and len(result[1]) > 10]
if valid:
best_r = self._judge_best_response(results, primary_purpose)
else:
# Il pool primario è interamente in rate limit, errore auth o timeout:
# prova un solo profilo per ogni provider sano, in ordine di purpose/tier.
excluded = {provider.name for provider, _response, _latency in results}
fallback_pool = self._inter_provider_fallback_pool(primary_purpose, excluded)
fallback_results = []
for fallback in fallback_pool:
result = await self._fetch_one(fallback, messages, temperature, max_tokens)
fallback_results.append(result)
if not result[1].startswith("ERROR:") and len(result[1]) > 10:
_logger.info(
"[fleet] inter-provider fallback succeeded on %s/%s",
fallback.name,
fallback.profile,
)
best_r = result[1]
break
else:
failed_names = [provider.name for provider, _response, _latency in results + fallback_results]
raise ProviderUnavailableError(failed_names)
# S-CACHE-1: Popolamento cache asincrono
if not best_r.startswith("🔴"):
asyncio.create_task(set_cached_response(messages, best_r))
return best_r
def _judge_best_response(self, results: List[Tuple[ProviderConfig, str, float]], target_purpose: str) -> str:
valid = [(p, r, t) for p, r, t in results if not r.startswith("ERROR:") and len(r) > 10]
if not valid:
raise ProviderUnavailableError([p.name for p, _r, _t in results])
def score(item):
p, r, t = item
s = len(r)
# Bonus purpose target (Specializzazione)
if p.purpose == target_purpose: s += 2000
# Bonus latenza (Performance mobile-first)
if t < 2.0: s += 500
if "```" in r: s += 300
return s
_, best_r, _ = max(valid, key=score)
return best_r
async def stream_chat(self, messages: list, model: Optional[str] = None, temperature: float = 0.7, max_tokens: int = 4096) -> AsyncIterator[str]:
"""Streaming con failover immediato tra i provider configurati."""
if not self.providers: self.providers = self._load_providers()
# Nessun provider configurato: feedback immediato all'utente invece di
# cadere silenziosamente nel loop vuoto e dare un messaggio generico.
if not self.providers:
raise ProviderUnavailableError([])
# Un client di ruolo può contenere un solo provider specializzato.
# Dopo il suo primario, integra la flotta runtime non duplicata: un limite
# temporaneo di quel provider non deve rendere indisponibile l'intero task.
providers = list(self.providers)
try:
for fallback in self._load_providers():
if not any(
current.name == fallback.name
and current.base_url == fallback.base_url
for current in providers
):
providers.append(fallback)
except Exception as exc:
_logger.debug("Streaming fleet expansion skipped: %s", type(exc).__name__)
# Un profilo sano per provider: se l’intero pool primario è in rate
# limit, il fallback passa automaticamente al provider successivo.
providers = self._inter_provider_fallback_pool("stream", providers=providers)
attempted: list[str] = []
for provider in providers:
attempted.append(provider.name)
emitted = False
try:
client = self._client_for(provider)
stream = await asyncio.to_thread(
client.chat.completions.create,
model=provider.default_model,
messages=messages,
temperature=temperature,
max_tokens=max_tokens,
stream=True,
)
iterator = iter(stream)
while True:
chunk = await asyncio.to_thread(next, iterator, None)
if chunk is None:
break
if chunk.choices and chunk.choices[0].delta.content:
emitted = True
yield chunk.choices[0].delta.content
self._record_success(provider)
return
except Exception as e:
self._record_failure(provider, e)
_logger.warning(
"Streaming fallito su %s/%s (emitted=%s): %s",
provider.name, provider.profile, emitted, e,
)
# Retry solo prima del primo chunk: dopo output parziale un
# retry produrrebbe testo duplicato o una risposta incoerente.
if emitted:
raise
continue
raise ProviderUnavailableError(attempted)