fix: sécurisation des insertions de modèles via dictionnaire temporaire pour éviter les Duplicate Key

This commit is contained in:
xavier committed 2026-10-07 19:40:09 +02:00
1 parent 007e723e4e
commit fd015fd48f
1 file changed
+35 -26
+35 -26
View File
@@ -3,6 +3,7 @@ import httpx
from openai import AsyncOpenAI from openai import AsyncOpenAI
from fastapi import HTTPException from fastapi import HTTPException
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from sqlalchemy.exc import IntegrityError
from .models import SystemSettings, AIModel, FinancialLog from .models import SystemSettings, AIModel, FinancialLog
from datetime import datetime, timezone from datetime import datetime, timezone
@@ -37,7 +38,7 @@ async def sync_finances(db: Session, settings: SystemSettings):
update_finance_db(db, "OpenRouter", balance, usage) update_finance_db(db, "OpenRouter", balance, usage)
except Exception as e: print(f"Erreur Finance OR: {e}") except Exception as e: print(f"Erreur Finance OR: {e}")
# Groq (Gratuit Beta) # Groq
if settings.groq_api_key: update_finance_db(db, "Groq", 999.0, 0.0) if settings.groq_api_key: update_finance_db(db, "Groq", 999.0, 0.0)
# DeepSeek # DeepSeek
@@ -49,10 +50,10 @@ async def sync_finances(db: Session, settings: SystemSettings):
update_finance_db(db, "DeepSeek", float(infos.get("total_balance", 0)), 0.0) update_finance_db(db, "DeepSeek", float(infos.get("total_balance", 0)), 0.0)
except Exception: pass except Exception: pass
# Mistral AI (Pas de route standard simple pour le budget public, on mock) # Mistral AI
if settings.mistral_api_key: update_finance_db(db, "Mistral", 0.0, 0.0) if settings.mistral_api_key: update_finance_db(db, "Mistral", 0.0, 0.0)
# Gemini (Quota lié à GCP, pas de budget direct simple via API clé) # Gemini
if settings.gemini_api_key: update_finance_db(db, "Gemini", 0.0, 0.0) if settings.gemini_api_key: update_finance_db(db, "Gemini", 0.0, 0.0)
try: db.commit() try: db.commit()
@@ -69,7 +70,9 @@ def update_finance_db(db, provider, balance, usage):
async def sync_providers_models(db: Session, settings: SystemSettings, sync_type: str = "Automatique"): async def sync_providers_models(db: Session, settings: SystemSettings, sync_type: str = "Automatique"):
added = 0 added = 0
# TIMEOUT PASSÉ À 90 SECONDES # On stocke les modèles en mémoire avant de les envoyer en base pour éviter les doublons
models_to_process = {}
async with httpx.AsyncClient(timeout=90.0) as client: async with httpx.AsyncClient(timeout=90.0) as client:
# 1. OpenRouter # 1. OpenRouter
if settings.openrouter_api_key: if settings.openrouter_api_key:
@@ -78,6 +81,8 @@ async def sync_providers_models(db: Session, settings: SystemSettings, sync_type
if resp.status_code == 200: if resp.status_code == 200:
for item in resp.json().get("data", []): for item in resp.json().get("data", []):
try: try:
m_id = item.get("id")
if not m_id: continue
pricing = item.get("pricing") or {} pricing = item.get("pricing") or {}
try: pp = float(pricing.get("prompt") or 0.0) * 1000000 try: pp = float(pricing.get("prompt") or 0.0) * 1000000
except: pp = 0.0 except: pp = 0.0
@@ -85,8 +90,7 @@ async def sync_providers_models(db: Session, settings: SystemSettings, sync_type
except: pc = 0.0 except: pc = 0.0
is_free = (pp == 0.0 and pc == 0.0) is_free = (pp == 0.0 and pc == 0.0)
desc = item.get("description", "Modèle OpenRouter.")[:200] + "..." desc = item.get("description", "Modèle OpenRouter.")[:200] + "..."
process_model(db, "openrouter", item.get("id", "inconnu"), item.get("name", "Inconnu"), desc, determine_domain(item.get("id", "")), is_free, item.get("context_length", 0), pp, pc) models_to_process[m_id] = {"provider": "openrouter", "name": item.get("name", "Inconnu"), "desc": desc, "domain": determine_domain(m_id), "is_free": is_free, "ctx": item.get("context_length", 0), "pp": pp, "pc": pc}
added += 1
except: pass except: pass
except Exception as e: print(f"Erreur OR Models: {e}") except Exception as e: print(f"Erreur OR Models: {e}")
@@ -97,8 +101,8 @@ async def sync_providers_models(db: Session, settings: SystemSettings, sync_type
if resp.status_code == 200: if resp.status_code == 200:
for item in resp.json().get("data", []): for item in resp.json().get("data", []):
try: try:
process_model(db, "groq", item["id"], item["id"].capitalize(), "Modèle rapide LPU Groq.", determine_domain(item["id"]), True, 8192, 0.0, 0.0) m_id = item["id"]
added += 1 models_to_process[m_id] = {"provider": "groq", "name": m_id.capitalize(), "desc": "Modèle rapide LPU Groq.", "domain": determine_domain(m_id), "is_free": True, "ctx": 8192, "pp": 0.0, "pc": 0.0}
except: pass except: pass
except Exception: pass except Exception: pass
@@ -109,8 +113,8 @@ async def sync_providers_models(db: Session, settings: SystemSettings, sync_type
if resp.status_code == 200: if resp.status_code == 200:
for item in resp.json().get("data", []): for item in resp.json().get("data", []):
try: try:
process_model(db, "deepseek", item["id"], item["id"].capitalize(), "Modèle officiel DeepSeek.", determine_domain(item["id"]), False, 64000, 0.14, 0.28) m_id = item["id"]
added += 1 models_to_process[m_id] = {"provider": "deepseek", "name": m_id.capitalize(), "desc": "Modèle officiel DeepSeek.", "domain": determine_domain(m_id), "is_free": False, "ctx": 64000, "pp": 0.14, "pc": 0.28}
except: pass except: pass
except Exception: pass except Exception: pass
@@ -121,8 +125,8 @@ async def sync_providers_models(db: Session, settings: SystemSettings, sync_type
if resp.status_code == 200: if resp.status_code == 200:
for item in resp.json().get("data", []): for item in resp.json().get("data", []):
try: try:
process_model(db, "mistral", item["id"], item["id"].capitalize(), "Modèle officiel Mistral AI.", determine_domain(item["id"]), False, 32000, 0.2, 0.6) m_id = item["id"]
added += 1 models_to_process[m_id] = {"provider": "mistral", "name": m_id.capitalize(), "desc": "Modèle officiel Mistral AI.", "domain": determine_domain(m_id), "is_free": False, "ctx": 32000, "pp": 0.2, "pc": 0.6}
except: pass except: pass
except Exception: pass except Exception: pass
@@ -135,9 +139,7 @@ async def sync_providers_models(db: Session, settings: SystemSettings, sync_type
try: try:
m_id = item["name"].replace("models/", "") m_id = item["name"].replace("models/", "")
desc = item.get("description", "Modèle Google Gemini.")[:200] + "..." desc = item.get("description", "Modèle Google Gemini.")[:200] + "..."
ctx = item.get("inputTokenLimit", 32000) models_to_process[m_id] = {"provider": "gemini", "name": item.get("displayName", m_id), "desc": desc, "domain": determine_domain(m_id), "is_free": True, "ctx": item.get("inputTokenLimit", 32000), "pp": 0.0, "pc": 0.0}
process_model(db, "gemini", m_id, item.get("displayName", m_id), desc, determine_domain(m_id), True, ctx, 0.0, 0.0)
added += 1
except: pass except: pass
except Exception: pass except Exception: pass
@@ -151,11 +153,27 @@ async def sync_providers_models(db: Session, settings: SystemSettings, sync_type
try: try:
m_id = item.get("name") m_id = item.get("name")
desc = item.get("description", "Modèle Cloudflare Workers AI.")[:200] + "..." desc = item.get("description", "Modèle Cloudflare Workers AI.")[:200] + "..."
process_model(db, "cloudflare", m_id, m_id.split("/")[-1], desc, determine_domain(m_id), True, 4096, 0.0, 0.0) models_to_process[m_id] = {"provider": "cloudflare", "name": m_id.split("/")[-1], "desc": desc, "domain": determine_domain(m_id), "is_free": True, "ctx": 4096, "pp": 0.0, "pc": 0.0}
added += 1
except: pass except: pass
except Exception: pass except Exception: pass
# --- Phase d'enregistrement sécurisée ---
for m_id, data in models_to_process.items():
try:
existing = db.query(AIModel).filter(AIModel.model_id == m_id).first()
if existing:
existing.pricing_prompt = data["pp"]; existing.pricing_completion = data["pc"]; existing.is_free = data["is_free"]; existing.last_updated = datetime.now(timezone.utc)
else:
db.add(AIModel(provider=data["provider"], model_id=m_id, name=data["name"], description_fr=data["desc"], domain=data["domain"], is_free=data["is_free"], context_length=data["ctx"], pricing_prompt=data["pp"], pricing_completion=data["pc"]))
added += 1
# Commit très fréquent pour éviter les gros blocs qui plantent
if added % 50 == 0:
db.commit()
except IntegrityError:
db.rollback() # Si conflit, on annule cette insertion et on continue
except Exception:
db.rollback()
try: try:
settings.last_sync_date = datetime.now(timezone.utc) settings.last_sync_date = datetime.now(timezone.utc)
settings.last_sync_type = sync_type settings.last_sync_type = sync_type
@@ -167,15 +185,6 @@ async def sync_providers_models(db: Session, settings: SystemSettings, sync_type
await sync_finances(db, settings) await sync_finances(db, settings)
return {"status": "success", "models_processed": added} return {"status": "success", "models_processed": added}
def process_model(db, provider, mod_id, name, desc, domain, is_free, ctx, pp, pc):
try:
existing = db.query(AIModel).filter(AIModel.model_id == mod_id).first()
if existing:
existing.pricing_prompt = pp; existing.pricing_completion = pc; existing.is_free = is_free; existing.last_updated = datetime.now(timezone.utc)
else:
db.add(AIModel(provider=provider, model_id=mod_id, name=name, description_fr=desc, domain=domain, is_free=is_free, context_length=ctx, pricing_prompt=pp, pricing_completion=pc))
except Exception: pass
def get_client_for_model(db: Session, model_id: str, settings: SystemSettings): def get_client_for_model(db: Session, model_id: str, settings: SystemSettings):
"""Récupère dynamiquement le bon client OpenAI en fonction du fournisseur du modèle.""" """Récupère dynamiquement le bon client OpenAI en fonction du fournisseur du modèle."""
model_db = db.query(AIModel).filter(AIModel.model_id == model_id).first() model_db = db.query(AIModel).filter(AIModel.model_id == model_id).first()