File size: 17,006 Bytes
7532f48
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9ab557a
7532f48
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9ab557a
 
 
 
 
 
 
 
7532f48
 
 
 
 
 
9ab557a
7532f48
 
 
 
 
 
 
 
 
9ab557a
 
 
 
7532f48
 
 
 
9ab557a
 
 
 
 
 
 
 
 
 
 
7532f48
9ab557a
7532f48
9ab557a
 
 
 
7532f48
 
9ab557a
7532f48
 
9ab557a
7532f48
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
import json
import os
import re

# Cache HF explicite, à la racine du projet (portable local/HF Spaces — le
# disque du Space est accessible en écriture par défaut sur le tier gratuit,
# pas besoin de stockage persistant payant pour ça). Ne PAS forcer
# HF_HUB_OFFLINE=1 : ça bloquerait le tout premier déploiement (cache vide,
# mode offline = téléchargement impossible = le Space ne démarre jamais).
os.environ.setdefault("HF_HOME", os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), ".hf_cache"))

import faiss
import numpy as np
import torch
from dotenv import load_dotenv
from openai import APIConnectionError, APITimeoutError, OpenAI
from rank_bm25 import BM25Okapi
from sentence_transformers import SentenceTransformer

load_dotenv()

INDEX_PATH = "faiss_solva2/index.faiss"
MAPPING_PATH = "faiss_solva2/mapping.json"
EMBEDDING_MODEL_NAME = "Shitao/bge-m3"  # miroir safetensors du même modèle que BAAI/bge-m3

# Fournisseurs LLM disponibles, tous compatibles OpenAI SDK (base_url + clé +
# nom de modèle changent, le reste du code est identique). Ne pas supprimer
# une entrée quand on change de défaut : ça permet de rebasculer en une ligne
# (provider="mistral"|"gemini"|...).
PROVIDERS = {
    "gemini": {
        "base_url": "https://generativelanguage.googleapis.com/v1beta/openai/",
        "api_key_env": "GEMINI_API_KEY",
        "model": "gemini-flash-latest",  # "gemini-2.5-flash" renvoie 404 (fermé aux nouveaux comptes)
    },
    "mistral": {
        "base_url": "https://api.mistral.ai/v1",
        "api_key_env": "MISTRAL_API_KEY",
        "model": "mistral-large-latest",
    },
    "groq": {
        "base_url": "https://api.groq.com/openai/v1",
        "api_key_env": "GROQ_API_KEY",
        "model": "llama-3.3-70b-versatile",
    },
    "cerebras": {
        "base_url": "https://api.cerebras.ai/v1",
        "api_key_env": "CEREBRAS_API_KEY",
        "model": "gpt-oss-120b",  # llama-3.3-70b indisponible sur ce compte (404)
    },
}

# Gemini retenu (Décision 008) : 1500 requêtes/jour, largement suffisant pour
# le pipeline conversationnel (3 appels LLM/question : reformulation +
# clarification + génération) qui saturait le tier gratuit Mistral.
DEFAULT_PROVIDER = "gemini"

# Décision 009 (Gemini->Mistral) puis Décision 012 (ajout de Groq en 3e,
# nouvelle clé validée) : chaîne de fallback automatique. Cerebras disponible
# dans PROVIDERS mais pas dans cette chaîne (pas testé en usage courant, pas
# de raison de l'y ajouter pour l'instant).
FALLBACK_ORDER = ["gemini", "mistral", "groq"]

_llm_clients = {}  # un client OpenAI par provider, créé à la demande


def _get_llm_client(provider):
    if provider not in _llm_clients:
        cfg = PROVIDERS[provider]
        _llm_clients[provider] = OpenAI(
            api_key=os.environ[cfg["api_key_env"]],
            base_url=cfg["base_url"],
            timeout=30,  # évite un blocage indéfini si l'API ne répond pas (défaut SDK ~10 min)
        )
    return _llm_clients[provider]


# Alias de compatibilité (utilisés par app.py et d'anciens scripts) : pointent
# vers le provider par défaut.
LLM_MODEL = PROVIDERS[DEFAULT_PROVIDER]["model"]
_llm_client = _get_llm_client(DEFAULT_PROVIDER)

# Gemini consomme une partie du budget max_tokens en raisonnement interne caché,
# même sur des prompts de classification triviaux (ex. "OUI"/"NON" a consommé
# ~150 tokens de réflexion avant le token de réponse visible). Pas de moyen
# documenté de désactiver ça via l'endpoint compatible OpenAI (reasoning_effort
# et extra_body google.thinking_config testés, tous deux rejetés en 400) — on
# compense en donnant assez de marge. Mistral et Groq n'ont pas ce
# comportement, donc pas besoin d'autant de marge de ce côté.
MAX_TOKENS_COURT = {"gemini": 2000, "mistral": 200, "groq": 200}     # classification (juge OUI/NON, CLAIRE/AMBIGUE)
MAX_TOKENS_LONG = {"gemini": 4096, "mistral": 2048, "groq": 2048}    # réponses substantielles (génération, reformulation)


# Codes déclenchant un fallback : quota/débit (429), facturation (402),
# authentification/clé invalide (401, 403 — et 400 : constaté empiriquement
# que l'endpoint Gemini répond 400 "Please pass a valid API key" pour une clé
# invalide, pas 401 comme la plupart des API OpenAI-compatibles).
CODES_FALLBACK = {400, 401, 402, 403, 429}


def appel_llm(prompt, taille="long", verbose=True, providers=None):
    """Appel LLM unifié avec fallback automatique (Décision 009).

    Essaie les providers de `providers` (par défaut FALLBACK_ORDER) dans
    l'ordre. Bascule sur le suivant en cas d'erreur de quota/débit/
    authentification (CODES_FALLBACK) OU d'erreur de connexion réseau
    (APIConnectionError/APITimeoutError — celles-ci n'ont PAS de status_code
    car l'échec survient avant toute réponse HTTP, donc un simple test sur
    status_code les ratait). Si tous échouent, ou sur toute autre erreur non
    couverte, l'exception remonte telle quelle — pas de 3e fallback, pas de
    boucle. Logs toujours en print(..., flush=True) pour apparaître
    immédiatement dans les logs du Space (pas de bufferisation).
    """
    providers = providers or FALLBACK_ORDER
    max_tokens_par_taille = MAX_TOKENS_COURT if taille == "court" else MAX_TOKENS_LONG

    derniere_erreur = None
    for i, provider in enumerate(providers):
        print(f"[appel_llm] tentative {i + 1}/{len(providers)} : provider={provider}", flush=True)
        try:
            client = _get_llm_client(provider)
            model_name = PROVIDERS[provider]["model"]
            completion = client.chat.completions.create(
                model=model_name,
                messages=[{"role": "user", "content": prompt}],
                temperature=0,
                max_tokens=max_tokens_par_taille[provider],
            )
            if i > 0:
                print(f"[appel_llm] SUCCÈS après bascule : {providers[i - 1]} -> {provider}", flush=True)
            elif verbose:
                print(f"[appel_llm] SUCCÈS sur {provider} (premier essai)", flush=True)
            return completion.choices[0].message.content
        except Exception as e:
            derniere_erreur = e
            status_code = getattr(e, "status_code", None)
            erreur_connexion = isinstance(e, (APIConnectionError, APITimeoutError))
            declenche_fallback = status_code in CODES_FALLBACK or erreur_connexion

            print(
                f"[appel_llm] ÉCHEC provider={provider} | "
                f"type={type(e).__name__} | status_code={status_code} | "
                f"erreur_connexion={erreur_connexion} | "
                f"message complet={e!r}",
                flush=True,
            )

            if not declenche_fallback:
                print(f"[appel_llm] ARRÊT : erreur non couverte par le fallback, remontée telle quelle", flush=True)
                raise
            if i == len(providers) - 1:
                print(f"[appel_llm] Dernier provider de la chaîne épuisé ({provider})", flush=True)
            else:
                print(f"[appel_llm] Bascule prévue vers : {providers[i + 1]}", flush=True)
            continue

    print(f"[appel_llm] TOUS LES PROVIDERS ONT ÉCHOUÉ : {providers}", flush=True)
    raise RuntimeError(
        f"Tous les fournisseurs LLM ont échoué ({', '.join(providers)}). "
        f"Dernière erreur : {derniere_erreur!r}"
    )


# Dictionnaire fixe (pas de LLM) : sigles courants de Solvabilité II.
# "SCR" apparaît tel quel dans le corpus (38 occurrences, ex. "capital de
# solvabilité requis (SCR)"), mais les autres sigles n'y figurent JAMAIS —
# le texte officiel les épelle toujours en toutes lettres. D'où l'intérêt
# de l'expansion : rapprocher la question posée (en jargon métier) du
# vocabulaire réellement utilisé dans le texte.
ACRONYMES = {
    "SCR": "capital de solvabilité requis",
    "MCR": "minimum de capital requis",
    "ORSA": "évaluation interne des risques et de la solvabilité",
    "BE": "meilleure estimation",
    "PT": "provisions techniques",
    "FP": "fonds propres",
    "VaR": "valeur à risque",
    "SFCR": "rapport sur la solvabilité et la situation financière",
    "EIOPA": "Autorité européenne des assurances et des pensions professionnelles",
    "AEAPP": "Autorité européenne des assurances et des pensions professionnelles",
    "ACPR": "Autorité de contrôle prudentiel et de résolution",
}


def expand_query(question):
    expanded = question
    for sigle, forme_longue in ACRONYMES.items():
        pattern = re.compile(rf"\b{re.escape(sigle)}\b", re.IGNORECASE)
        if pattern.search(expanded):
            expanded = pattern.sub(f"{sigle} ({forme_longue})", expanded, count=1)
    return expanded


PROMPT_TEMPLATE = """Tu es un assistant spécialisé sur la directive Solvabilité II
(2009/138/CE). Tu réponds UNIQUEMENT à partir des articles fournis
ci-dessous.
RÈGLES ABSOLUES :
1. Réponds uniquement d'après les articles fournis.
2. Cite TOUJOURS l'article source au format [Article N].
3. Si la question contient un chiffre, un pourcentage, un seuil ou une
   borne inexact(e), et que les articles fournis donnent la valeur
   correcte, NE REFUSE PAS : signale explicitement l'écart avec le
   chiffre mentionné dans la question, indique la valeur correcte, et
   cite l'article source. Ceci prime sur la règle 4. Si tu corriges une
   valeur, ne commence JAMAIS ta réponse par la phrase de refus.
   Commence directement par la correction, ex : "Le SCR est calibré à
   99,5 %, et non 99,9 % [Article X]."
4. Si l'information n'est pas dans les articles fournis, réponds
   exactement : "Je ne trouve pas cette information dans les
   articles fournis."
5. N'invente jamais. Ne complète jamais de mémoire.

ARTICLES FOURNIS :
{contexte}

QUESTION : {question}
RÉPONSE :"""

# --- Chargement une fois au chargement du module ---
_index = faiss.read_index(INDEX_PATH)
with open(MAPPING_PATH, encoding="utf-8") as _f:
    _mapping = json.load(_f)

_embedding_model = None  # chargement paresseux


def _get_embedding_model():
    global _embedding_model
    if _embedding_model is None:
        device = "cuda" if torch.cuda.is_available() else "cpu"
        _embedding_model = SentenceTransformer(EMBEDDING_MODEL_NAME, device=device)
        _embedding_model.max_seq_length = 1024
    return _embedding_model


def _rechercher(question, k):
    model = _get_embedding_model()
    q_emb = model.encode([question], normalize_embeddings=True, convert_to_numpy=True).astype(np.float32)
    _, idxs = _index.search(q_emb, k)
    return [_mapping[i] for i in idxs[0]]


def _rechercher_dense_idx(question, k):
    model = _get_embedding_model()
    q_emb = model.encode([question], normalize_embeddings=True, convert_to_numpy=True).astype(np.float32)
    _, idxs = _index.search(q_emb, k)
    return list(idxs[0])


# --- Retrieval hybride EN TEST (mode="hybrid" dans poser_question) : BM25 en
# complément du dense BGE-M3, ne remplace pas le retrieval actuel tant que non
# validé sur le golden set. ---
def _tokeniser_bm25(texte):
    return re.findall(r"\w+", texte.lower())


_bm25_corpus_tokens = [_tokeniser_bm25(f"{a['titre']} {a['texte']}") for a in _mapping]
_bm25 = BM25Okapi(_bm25_corpus_tokens)


def _rechercher_bm25_idx(question, k):
    scores = _bm25.get_scores(_tokeniser_bm25(question))
    return list(np.argsort(scores)[::-1][:k])


def retrieval_hybride(question, k=5, k_fusion=20, poids_dense=0.7, poids_bm25=0.3, rrf_k=60):
    """Fusion par Reciprocal Rank Fusion (RRF) du dense (BGE-M3/FAISS) et de
    BM25, pondérée 70% dense / 30% BM25 par défaut."""
    dense_idx = _rechercher_dense_idx(question, k_fusion)
    bm25_idx = _rechercher_bm25_idx(question, k_fusion)

    scores_fusion = {}
    for rang, idx in enumerate(dense_idx, start=1):
        idx = int(idx)
        scores_fusion[idx] = scores_fusion.get(idx, 0.0) + poids_dense * (1 / (rrf_k + rang))
    for rang, idx in enumerate(bm25_idx, start=1):
        idx = int(idx)
        scores_fusion[idx] = scores_fusion.get(idx, 0.0) + poids_bm25 * (1 / (rrf_k + rang))

    tries = sorted(scores_fusion.items(), key=lambda x: x[1], reverse=True)[:k]
    return [_mapping[idx] for idx, _ in tries]


def _construire_contexte(articles):
    blocs = []
    for a in articles:
        blocs.append(f"[Article {a['numero_article']}{a['titre']}]\n{a['texte']}\n")
    return "\n".join(blocs)


def _est_un_refus(reponse):
    return reponse.strip().startswith("Je ne trouve pas") and "[Article" not in reponse


JUGE_PROMPT_TEMPLATE = """Voici une question et des articles de loi récupérés pour y répondre.

QUESTION : {question}

ARTICLES RÉCUPÉRÉS :
{contexte}

Ces articles contiennent-ils de quoi répondre à la question ?
Réponds STRICTEMENT sous l'une de ces deux formes, rien d'autre :
OUI
NON : <reformulation courte de la requête de recherche>"""


def _juger_retrieval(question, articles, verbose=True):
    contexte = _construire_contexte(articles)
    prompt = JUGE_PROMPT_TEMPLATE.format(question=question, contexte=contexte)

    reponse = appel_llm(prompt, taille="court", verbose=verbose).strip()

    if reponse.upper().startswith("OUI"):
        return True, None
    if reponse.upper().startswith("NON"):
        partie = reponse.split(":", 1)
        reformulation = partie[1].strip() if len(partie) > 1 else None
        return False, reformulation
    # réponse hors format attendu : par prudence, on considère le retrieval suffisant
    # (évite de déclencher une reformulation sur une base non fiable)
    return True, None


CLARIFICATION_JUGE_PROMPT = """Cette question sur Solvabilité II est-elle assez précise pour y répondre, ou trop vague/ambiguë ?
Réponds uniquement par un seul mot : CLAIRE ou AMBIGUE.

QUESTION : {question}"""

CLARIFICATION_REFORMULATION_PROMPT = """Cette question sur Solvabilité II est trop vague ou ambiguë pour y répondre directement :

QUESTION : {question}

Propose UNE seule question de clarification précise, en français, pour aider
l'utilisateur à préciser sa demande. Réponds uniquement avec cette question
de clarification, rien d'autre."""


def clarifier(question, verbose=True):
    decision = appel_llm(
        CLARIFICATION_JUGE_PROMPT.format(question=question), taille="court", verbose=verbose
    ).strip().upper()

    if "AMBIG" not in decision:
        return {"ambigue": False, "message": None}

    message = appel_llm(
        CLARIFICATION_REFORMULATION_PROMPT.format(question=question), taille="long", verbose=verbose
    ).strip()
    return {"ambigue": True, "message": message}


def poser_question(question, k=5, verbose=True, self_eval=False, clarify=False, mode="dense"):
    if clarify:
        resultat_clarif = clarifier(question, verbose=verbose)
        if resultat_clarif["ambigue"]:
            if verbose:
                print(f"QUESTION : {question}")
                print(f"[clarification] question jugée AMBIGUË — pas de retrieval, pas de génération")
                print(f"CLARIFICATION DEMANDÉE : {resultat_clarif['message']}")
                print()
            return {"type": "clarification", "message": resultat_clarif["message"]}

    requete_recherche = expand_query(question)
    articles = retrieval_hybride(requete_recherche, k) if mode == "hybrid" else _rechercher(requete_recherche, k)

    if self_eval:
        suffisant, reformulation = _juger_retrieval(question, articles, verbose=verbose)
        if not suffisant and reformulation:
            if verbose:
                print(f"[self-eval] retrieval jugé insuffisant, reformulation : \"{reformulation}\"")
            reformulation_expandue = expand_query(reformulation)
            articles = (
                retrieval_hybride(reformulation_expandue, k)
                if mode == "hybrid"
                else _rechercher(reformulation_expandue, k)
            )
        elif not suffisant and verbose:
            print("[self-eval] retrieval jugé insuffisant, mais pas de reformulation exploitable — on garde le résultat initial")

    contexte = _construire_contexte(articles)
    prompt = PROMPT_TEMPLATE.format(contexte=contexte, question=question)

    reponse = appel_llm(prompt, taille="long", verbose=verbose)

    if verbose:
        print(f"QUESTION : {question}")
        print(f"RÉPONSE  : {reponse}")
        refus = _est_un_refus(reponse)
        print(f"REFUS DÉTECTÉ : {refus}")
        print("SOURCES CONSULTÉES :")
        for a in articles:
            print(f"  - Article {a['numero_article']}{a['titre']}")
        print()

    return reponse


if __name__ == "__main__":
    questions_test = [
        "Comment calcule-t-on le minimum de capital requis ?",
        "Comment sont calculées les provisions techniques ?",
        "Quel est le taux de TVA sur les croissants ?",
    ]

    for q in questions_test:
        poser_question(q)
        print("-" * 70)