fix: sécurisation des insertions de modèles via dictionnaire temporaire pour éviter les Duplicate Key
This commit is contained in:
1 parent
007e723e4e
commit
fd015fd48f
1 file changed
+35
-26
+35
-26
@@ -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()
|
||||||
|
|||||||
Reference in new issue
Block a user