Plusieurs gros ajouts :
All checks were successful
All checks were successful
- Possibilité de discuter avec un PDF ; RAG ou analyse approfondie. Enlèvement de l'autre outil PDF de discussion qui analysait d'abord un PDF en proposant directement une intégration sans attendre qu'on pose de question - Mise en place de l'import directement dans les outils dans la sidebar - Mise en place d'un outil pour créer des tables aléatoires avec possibilité d'utiliser pendant la partie - Mise en place d'un outil pour mettre en place des PNJ, scènes, chapitre.... directement à partir de la discussion avec le PDF - Mise en place RAG avec mistal-embeding ou nomic si on utilise ollama - Mise en place mistral, google en fournisseurs alternatifs pour l'IA dans le cloud - version 0.11.0-bêta
This commit is contained in:
@@ -53,3 +53,27 @@ def _split_oversized(paragraph: str, enc, target_tokens: int) -> list[str]:
|
||||
for i in range(0, len(tokens), target_tokens):
|
||||
out.append(enc.decode(tokens[i : i + target_tokens]))
|
||||
return out
|
||||
|
||||
|
||||
def split_in_half(text: str) -> tuple[str, str]:
|
||||
"""Coupe `text` en deux moitiés ~égales, de préférence sur un saut de ligne
|
||||
proche du milieu (pour ne pas trancher en plein mot/phrase).
|
||||
|
||||
Sert au repli anti-troncature des imports : quand la SORTIE d'un morceau est
|
||||
coupée (le modèle ne peut pas tout réécrire en une réponse), on retraite ce
|
||||
morceau en deux moitiés. Renvoie ('', '') si le texte est trop court pour
|
||||
être découpé utilement (garde-fou anti-récursion infinie).
|
||||
"""
|
||||
text = text.strip()
|
||||
if len(text) < 400:
|
||||
return "", ""
|
||||
mid = len(text) // 2
|
||||
# Cherche un saut de ligne juste avant le milieu, sinon juste après.
|
||||
cut = text.rfind("\n", 0, mid)
|
||||
if cut < len(text) // 4:
|
||||
nxt = text.find("\n", mid)
|
||||
cut = nxt if nxt != -1 else mid
|
||||
left, right = text[:cut].strip(), text[cut:].strip()
|
||||
if not left or not right:
|
||||
return "", ""
|
||||
return left, right
|
||||
|
||||
20
brain/app/application/embeddings.py
Normal file
20
brain/app/application/embeddings.py
Normal file
@@ -0,0 +1,20 @@
|
||||
"""Port d'embeddings (RAG des notebooks).
|
||||
|
||||
Abstraction du calcul de vecteurs : un texte → une liste de floats. Les adapters
|
||||
concrets (Ollama local, Mistral cloud) la satisfont par duck typing, comme pour
|
||||
les LLMProvider. Le RAG n'en dépend que via cette interface.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
|
||||
class EmbeddingError(Exception):
|
||||
"""Échec du calcul d'embeddings (modèle indisponible, réseau, quota…)."""
|
||||
|
||||
|
||||
class EmbeddingProvider(Protocol):
|
||||
"""Calcule les vecteurs d'une liste de textes (ordre préservé)."""
|
||||
|
||||
async def embed(self, texts: list[str]) -> list[list[float]]:
|
||||
...
|
||||
@@ -12,9 +12,14 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from app.application.chunking import chunk_text
|
||||
from app.application.llm_json import load_json_object
|
||||
from app.application.chunking import chunk_text, split_in_half
|
||||
from app.application.llm_json import load_json_object, looks_like_truncated_json
|
||||
from app.application.llm_retry import generate_with_retry
|
||||
from app.application.streaming import with_heartbeat
|
||||
|
||||
# Repli anti-troncature : si la sortie d'un morceau est coupée, on le retraite en
|
||||
# 2 moitiés. Borné en profondeur (3 niveaux => jusqu'à 8 sous-blocs).
|
||||
_MAX_SPLIT_DEPTH = 3
|
||||
from app.domain.models import (
|
||||
ArcProposal,
|
||||
CampaignImportResult,
|
||||
@@ -22,7 +27,7 @@ from app.domain.models import (
|
||||
RoomProposal,
|
||||
SceneProposal,
|
||||
)
|
||||
from app.domain.ports import LLMProvider, PdfTextExtractor
|
||||
from app.domain.ports import LLMProvider, LLMProviderError, PdfTextExtractor
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -229,8 +234,30 @@ class ImportCampaignUseCase:
|
||||
}
|
||||
|
||||
merger = _TreeMerger()
|
||||
skipped = 0
|
||||
last_error: str | None = None
|
||||
for i, chunk in enumerate(chunks):
|
||||
merger.add(await self._map_chunk(chunk, index=i, total=total))
|
||||
# RÉSILIENCE : un morceau qui échoue (provider saturé, quota, etc.) est
|
||||
# SAUTÉ — on ne perd pas tout l'import pour autant. On n'abandonne que
|
||||
# si AUCUN morceau ne passe (cf. après la boucle).
|
||||
# HEARTBEAT : keep-alive pendant l'appel LLM pour ne jamais laisser le
|
||||
# flux SSE silencieux (sinon le Core coupe sur timeout d'inactivité).
|
||||
try:
|
||||
arcs_payload: list[dict] | None = None
|
||||
async for kind, payload in with_heartbeat(
|
||||
self._map_chunk(chunk, index=i, total=total)
|
||||
):
|
||||
if kind == "heartbeat":
|
||||
yield {"type": "heartbeat", "current": i + 1, "total": total}
|
||||
else:
|
||||
arcs_payload = payload
|
||||
merger.add(arcs_payload or [])
|
||||
except LLMProviderError as exc:
|
||||
skipped += 1
|
||||
last_error = str(exc)
|
||||
logger.warning("Morceau %s/%s ignoré (échec LLM) : %s", i + 1, total, exc)
|
||||
yield {"type": "chunk_failed", "current": i + 1, "total": total,
|
||||
"message": str(exc)[:300]}
|
||||
arcs, chapters, scenes = merger.counts()
|
||||
yield {
|
||||
"type": "progress",
|
||||
@@ -239,42 +266,74 @@ class ImportCampaignUseCase:
|
||||
"arc_count": arcs,
|
||||
"chapter_count": chapters,
|
||||
"scene_count": scenes,
|
||||
"skipped": skipped,
|
||||
}
|
||||
|
||||
if total > 0 and skipped == total:
|
||||
# Tout a échoué : "done" vide serait trompeur → erreur explicite.
|
||||
yield {"type": "error",
|
||||
"message": "Tous les morceaux ont échoué auprès du fournisseur IA. "
|
||||
f"Dernier message : {last_error or 'inconnu'}"}
|
||||
return
|
||||
|
||||
yield {
|
||||
"type": "done",
|
||||
"arcs": _serialize_arcs(merger.result()),
|
||||
"page_count": doc.page_count,
|
||||
"ocr_page_count": doc.ocr_page_count,
|
||||
"skipped": skipped,
|
||||
}
|
||||
|
||||
# --- MAP : un morceau → sous-arbre ---------------------------------------
|
||||
|
||||
async def _map_chunk(self, chunk: str, *, index: int, total: int) -> list[dict]:
|
||||
return await self._extract_arcs(chunk, index=index, total=total, depth=0)
|
||||
|
||||
async def _extract_arcs(
|
||||
self, text: str, *, index: int, total: int, depth: int
|
||||
) -> list[dict]:
|
||||
"""Extrait l'arborescence d'un texte. Si la SORTIE est tronquée, retraite le
|
||||
texte en DEUX moitiés et concatène — le `_TreeMerger` final dédoublonne par
|
||||
nom (un arc/chapitre coupé entre les moitiés est recollé)."""
|
||||
prompt = (
|
||||
_MAP_SYSTEM.format(default_arc=_DEFAULT_ARC_NAME)
|
||||
+ f"\n\n--- EXTRAIT {index + 1}/{total} ---\n{chunk}\n\n"
|
||||
+ f"\n\n--- EXTRAIT {index + 1}/{total} ---\n{text}\n\n"
|
||||
"Renvoie maintenant le JSON de l'arborescence."
|
||||
)
|
||||
raw = await generate_with_retry(
|
||||
self._llm, prompt, output_format="json", temperature=_TEMPERATURE)
|
||||
return self._parse_arcs(raw, index=index)
|
||||
arcs, truncated = self._parse_arcs(raw, index=index)
|
||||
|
||||
if truncated and depth < _MAX_SPLIT_DEPTH:
|
||||
left, right = split_in_half(text)
|
||||
if left and right:
|
||||
logger.info(
|
||||
"Morceau %s : sortie tronquée → re-découpage en 2 moitiés (niveau %s).",
|
||||
index, depth + 1)
|
||||
a = await self._extract_arcs(left, index=index, total=total, depth=depth + 1)
|
||||
b = await self._extract_arcs(right, index=index, total=total, depth=depth + 1)
|
||||
return a + b
|
||||
if truncated:
|
||||
logger.warning(
|
||||
"Morceau %s : sortie tronquée, profondeur max atteinte — partiel conservé.", index)
|
||||
return arcs
|
||||
|
||||
@staticmethod
|
||||
def _parse_arcs(raw: str, *, index: int) -> list[dict]:
|
||||
"""Parse robuste : objet JSON équilibré, ou récupération partielle si tronqué."""
|
||||
def _parse_arcs(raw: str, *, index: int) -> tuple[list[dict], bool]:
|
||||
"""Parse robuste → (arcs, tronqué). `tronqué`=True si récupération partielle."""
|
||||
parsed, recovered = load_json_object(raw)
|
||||
if parsed is None:
|
||||
logger.warning("Morceau %s : aucun objet JSON exploitable, ignoré.", index)
|
||||
return []
|
||||
if recovered:
|
||||
logger.warning(
|
||||
"Morceau %s : sortie tronquée — récupération des éléments complets "
|
||||
"(envisagez des morceaux plus petits).", index)
|
||||
truncated = looks_like_truncated_json(raw)
|
||||
if not truncated:
|
||||
logger.warning(
|
||||
"Morceau %s : aucun objet JSON exploitable, ignoré. "
|
||||
"Début de la réponse du modèle : %r",
|
||||
index, (raw or "").strip()[:300] or "(réponse VIDE)")
|
||||
return [], truncated
|
||||
if isinstance(parsed, dict):
|
||||
arcs = parsed.get("arcs", [])
|
||||
return arcs if isinstance(arcs, list) else []
|
||||
return []
|
||||
return (arcs if isinstance(arcs, list) else []), recovered
|
||||
return [], recovered
|
||||
|
||||
|
||||
def _serialize_arcs(arcs: list[ArcProposal]) -> list[dict]:
|
||||
|
||||
@@ -14,9 +14,16 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from app.application.chunking import CHUNK_TARGET_TOKENS, chunk_text
|
||||
from app.application.llm_json import load_json_object
|
||||
from app.application.chunking import CHUNK_TARGET_TOKENS, chunk_text, split_in_half
|
||||
from app.application.llm_json import load_json_object, looks_like_truncated_json
|
||||
from app.application.llm_retry import generate_with_retry
|
||||
from app.application.streaming import with_heartbeat
|
||||
|
||||
# Repli anti-troncature : si la SORTIE d'un morceau est coupée (le modèle ne peut
|
||||
# pas tout réécrire en une réponse), on retraite ce morceau en 2 moitiés. Borné en
|
||||
# profondeur pour éviter une récursion infinie (3 niveaux => jusqu'à 8 sous-blocs ;
|
||||
# 1-2 niveaux suffisent en pratique, le reste est un garde-fou).
|
||||
_MAX_SPLIT_DEPTH = 3
|
||||
from app.domain.models import RulesImportResult
|
||||
from app.domain.ports import LLMProvider, LLMProviderError, PdfTextExtractor
|
||||
|
||||
@@ -95,6 +102,24 @@ class _SectionMerger:
|
||||
return {title: "\n\n".join(parts) for title, parts in self._merged.items()}
|
||||
|
||||
|
||||
def _combine_sections(a: dict[str, str], b: dict[str, str]) -> dict[str, str]:
|
||||
"""Fusionne deux dicts de sections (issus des 2 moitiés d'un morceau re-découpé).
|
||||
|
||||
Titres insensibles à la casse : un même titre présent des deux côtés (une section
|
||||
coupée par le re-découpage) voit ses contenus concaténés au lieu d'être écrasés.
|
||||
"""
|
||||
out = dict(a)
|
||||
by_lower = {k.lower(): k for k in out}
|
||||
for title, content in b.items():
|
||||
key = by_lower.get(title.lower())
|
||||
if key is not None:
|
||||
out[key] = f"{out[key]}\n\n{content}".strip()
|
||||
else:
|
||||
out[title] = content
|
||||
by_lower[title.lower()] = title
|
||||
return out
|
||||
|
||||
|
||||
class ImportRulesUseCase:
|
||||
"""Transforme un PDF de règles en proposition de sections markdown."""
|
||||
|
||||
@@ -152,48 +177,103 @@ class ImportRulesUseCase:
|
||||
}
|
||||
|
||||
merger = _SectionMerger()
|
||||
skipped = 0
|
||||
last_error: str | None = None
|
||||
for i, chunk in enumerate(chunks):
|
||||
new_titles = merger.add(await self._map_chunk(chunk, index=i, total=total))
|
||||
# RÉSILIENCE : un morceau qui échoue est SAUTÉ, l'import continue.
|
||||
# Abandon seulement si AUCUN morceau ne passe (cf. après la boucle).
|
||||
# HEARTBEAT : on émet des keep-alive pendant l'appel LLM (long sur un
|
||||
# provider lent) pour que le flux SSE ne soit jamais coupé par le Core.
|
||||
new_titles: list[str] = []
|
||||
try:
|
||||
sections: dict[str, str] | None = None
|
||||
async for kind, payload in with_heartbeat(
|
||||
self._map_chunk(chunk, index=i, total=total)
|
||||
):
|
||||
if kind == "heartbeat":
|
||||
yield {"type": "heartbeat", "current": i + 1, "total": total}
|
||||
else:
|
||||
sections = payload
|
||||
new_titles = merger.add(sections or {})
|
||||
except LLMProviderError as exc:
|
||||
skipped += 1
|
||||
last_error = str(exc)
|
||||
logger.warning("Morceau %s/%s ignoré (échec LLM) : %s", i + 1, total, exc)
|
||||
yield {"type": "chunk_failed", "current": i + 1, "total": total,
|
||||
"message": str(exc)[:300]}
|
||||
yield {
|
||||
"type": "progress",
|
||||
"current": i + 1,
|
||||
"total": total,
|
||||
"new_sections": new_titles,
|
||||
"skipped": skipped,
|
||||
}
|
||||
|
||||
if total > 0 and skipped == total:
|
||||
yield {"type": "error",
|
||||
"message": "Tous les morceaux ont échoué auprès du fournisseur IA. "
|
||||
f"Dernier message : {last_error or 'inconnu'}"}
|
||||
return
|
||||
|
||||
yield {
|
||||
"type": "done",
|
||||
"sections": merger.result(),
|
||||
"page_count": doc.page_count,
|
||||
"ocr_page_count": doc.ocr_page_count,
|
||||
"skipped": skipped,
|
||||
}
|
||||
|
||||
# --- MAP : un morceau → sections -----------------------------------------
|
||||
|
||||
async def _map_chunk(self, chunk: str, *, index: int, total: int) -> dict[str, str]:
|
||||
return await self._extract_sections(chunk, index=index, total=total, depth=0)
|
||||
|
||||
async def _extract_sections(
|
||||
self, text: str, *, index: int, total: int, depth: int
|
||||
) -> dict[str, str]:
|
||||
"""Extrait les sections d'un texte. Si la SORTIE est tronquée, retraite le
|
||||
texte en DEUX moitiés (chacune produit une réponse complète) et fusionne —
|
||||
ainsi aucune section n'est perdue, quel que soit le plafond de sortie."""
|
||||
prompt = (
|
||||
_MAP_SYSTEM.format(
|
||||
canonical="\n".join(f" - {s}" for s in _CANONICAL_SECTIONS)
|
||||
)
|
||||
+ f"\n\n--- EXTRAIT {index + 1}/{total} ---\n{chunk}\n\n"
|
||||
+ f"\n\n--- EXTRAIT {index + 1}/{total} ---\n{text}\n\n"
|
||||
"Renvoie maintenant le JSON des sections."
|
||||
)
|
||||
raw = await generate_with_retry(
|
||||
self._llm, prompt, output_format="json", temperature=_TEMPERATURE)
|
||||
return self._parse_sections(raw, index=index)
|
||||
sections, truncated = self._parse_sections(raw, index=index)
|
||||
|
||||
if truncated and depth < _MAX_SPLIT_DEPTH:
|
||||
left, right = split_in_half(text)
|
||||
if left and right:
|
||||
logger.info(
|
||||
"Morceau %s : sortie tronquée → re-découpage en 2 moitiés (niveau %s).",
|
||||
index, depth + 1)
|
||||
a = await self._extract_sections(left, index=index, total=total, depth=depth + 1)
|
||||
b = await self._extract_sections(right, index=index, total=total, depth=depth + 1)
|
||||
return _combine_sections(a, b)
|
||||
if truncated:
|
||||
logger.warning(
|
||||
"Morceau %s : sortie tronquée, profondeur max atteinte — partiel conservé.", index)
|
||||
return sections
|
||||
|
||||
@staticmethod
|
||||
def _parse_sections(raw: str, *, index: int) -> dict[str, str]:
|
||||
"""Parse robuste : objet JSON équilibré, ou récupération partielle si tronqué."""
|
||||
def _parse_sections(raw: str, *, index: int) -> tuple[dict[str, str], bool]:
|
||||
"""Parse robuste → (sections, tronqué). `tronqué`=True si récupération partielle."""
|
||||
parsed, recovered = load_json_object(raw)
|
||||
if parsed is None:
|
||||
logger.warning("Morceau %s : aucun objet JSON exploitable, ignoré.", index)
|
||||
return {}
|
||||
if recovered:
|
||||
logger.warning(
|
||||
"Morceau %s : sortie tronquée — récupération des sections complètes "
|
||||
"(envisagez des morceaux plus petits).", index)
|
||||
# Rien d'exploitable : soit prose (échec), soit JSON coupé avant toute
|
||||
# structure complète (→ on signalera 'tronqué' pour re-découper).
|
||||
truncated = looks_like_truncated_json(raw)
|
||||
if not truncated:
|
||||
logger.warning(
|
||||
"Morceau %s : aucun objet JSON exploitable, ignoré. "
|
||||
"Début de la réponse du modèle : %r",
|
||||
index, (raw or "").strip()[:300] or "(réponse VIDE)")
|
||||
return {}, truncated
|
||||
if not isinstance(parsed, dict):
|
||||
logger.warning("Morceau %s : le LLM n'a pas renvoyé un objet, ignoré.", index)
|
||||
return {}
|
||||
return {str(k): str(v) for k, v in parsed.items()}
|
||||
return {}, False
|
||||
return {str(k): str(v) for k, v in parsed.items()}, recovered
|
||||
|
||||
@@ -13,6 +13,16 @@ et tout ce qui suit. Renvoie None si aucun objet complet n'est trouvé
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
|
||||
# Blocs de "réflexion" des modèles raisonneurs (Nemotron, DeepSeek-R1, QwQ…).
|
||||
# Leur contenu est de la prose truffée d'accolades qui piège le détecteur de JSON
|
||||
# (et n'est jamais la réponse) → on le retire avant toute analyse.
|
||||
_REASONING_RE = re.compile(r"<think(?:ing)?>.*?</think(?:ing)?>", re.DOTALL | re.IGNORECASE)
|
||||
|
||||
|
||||
def _strip_reasoning(raw: str) -> str:
|
||||
return _REASONING_RE.sub("", raw)
|
||||
|
||||
|
||||
def load_json_object(raw: str) -> tuple[object | None, bool]:
|
||||
@@ -24,6 +34,7 @@ def load_json_object(raw: str) -> tuple[object | None, bool]:
|
||||
auquel cas le second élément vaut True.
|
||||
(None, False) si rien d'exploitable.
|
||||
"""
|
||||
raw = _strip_reasoning(raw)
|
||||
obj = extract_json_object(raw)
|
||||
if obj is not None:
|
||||
try:
|
||||
@@ -39,6 +50,18 @@ def load_json_object(raw: str) -> tuple[object | None, bool]:
|
||||
return None, False
|
||||
|
||||
|
||||
def looks_like_truncated_json(raw: str) -> bool:
|
||||
"""La sortie ressemble-t-elle à un JSON COUPÉ (accolades/crochets non refermés)
|
||||
plutôt qu'à de la prose ? Sert à déclencher un re-découpage même quand RIEN n'a
|
||||
pu être récupéré (cas où le 1er contenu est si long qu'il est coupé avant toute
|
||||
sous-structure complète). On exige un contenu substantiel pour éviter les
|
||||
faux positifs sur une courte réponse non-JSON."""
|
||||
s = (raw or "").strip()
|
||||
if "{" not in s or len(s) < 100:
|
||||
return False
|
||||
return s.count("{") > s.count("}") or s.count("[") > s.count("]")
|
||||
|
||||
|
||||
def extract_json_object(raw: str) -> str | None:
|
||||
if not raw:
|
||||
return None
|
||||
|
||||
@@ -18,7 +18,10 @@ from app.domain.ports import LLMProvider, LLMProviderError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_ATTEMPTS = 4
|
||||
# 3 tentatives : assez pour absorber un hoquet transitoire, sans s'acharner des
|
||||
# minutes sur un modèle durablement lent/saturé (les heartbeats gardent le flux
|
||||
# vivant, mais inutile de faire patienter l'utilisateur 15 min pour rien).
|
||||
_ATTEMPTS = 3
|
||||
_BASE_DELAY_SECONDS = 3.0
|
||||
# Un rate limit (429) "par minute" ne se libère pas en 2-3s : on attend plus
|
||||
# longtemps pour ces erreurs-là (le free tier OpenRouter plafonne ~20 req/min).
|
||||
|
||||
90
brain/app/application/notebook_chat.py
Normal file
90
brain/app/application/notebook_chat.py
Normal file
@@ -0,0 +1,90 @@
|
||||
"""Use case : chat ANCRÉ sur les sources d'un notebook (RAG).
|
||||
|
||||
À chaque message, on retrouve les passages pertinents des sources (via le RAG) et
|
||||
on les injecte dans le prompt système, en plus du contexte de campagne. Le modèle
|
||||
répond donc en s'appuyant sur la/les source(s) — pas sur ses connaissances générales.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import AsyncIterator
|
||||
|
||||
from app.application.notebook_rag import NotebookRagUseCase
|
||||
from app.domain.models import ChatMessage
|
||||
from app.domain.ports import LLMChatProvider
|
||||
|
||||
_SYSTEM_PROMPT = """Tu es un assistant de jeu de rôle qui aide à ADAPTER une source (PDF) à la CAMPAGNE de l'utilisateur.
|
||||
|
||||
Tu disposes de DEUX connaissances, toutes deux ci-dessous :
|
||||
1) LA CAMPAGNE de l'utilisateur (sa structure arcs/chapitres/scènes, ses PNJ, son univers) ;
|
||||
2) LA SOURCE (extraits pertinents du PDF).
|
||||
|
||||
Règles :
|
||||
- Pour une question sur SA CAMPAGNE (ex. « mon chapitre 3 », « mes PNJ »), appuie-toi sur la section CAMPAGNE.
|
||||
- Pour une question sur le livre, appuie-toi sur les EXTRAITS DE LA SOURCE.
|
||||
- CROISE les deux pour proposer des adaptations cohérentes avec sa campagne existante.
|
||||
- N'invente pas ce qui ne figure ni dans la campagne ni dans la source ; si tu ne sais pas, dis-le.
|
||||
- Quand un extrait porte un numéro de page (« (p. 12) »), cite-le (« d'après la p. 12 »).
|
||||
|
||||
{context_block}
|
||||
--- EXTRAITS PERTINENTS DE LA SOURCE ---
|
||||
{sources_block}
|
||||
--- FIN DES EXTRAITS ---
|
||||
|
||||
PROPOSITIONS D'INTÉGRATION (IMPORTANT) :
|
||||
Quand l'utilisateur veut CRÉER ou ADAPTER un élément concret pour sa campagne (un PNJ,
|
||||
une scène, un chapitre, un arc, une table aléatoire), termine ta réponse par un ou
|
||||
plusieurs BLOCS D'ACTION — un objet JSON par bloc, dans une clôture ```loremind-action.
|
||||
L'interface les transformera en boutons « Créer dans la campagne ». N'en mets que si
|
||||
c'est pertinent et explicitement souhaité. Formats acceptés :
|
||||
|
||||
```loremind-action
|
||||
{{"type": "npc", "name": "Nom", "description": "Fiche en quelques phrases."}}
|
||||
```
|
||||
```loremind-action
|
||||
{{"type": "scene", "name": "Nom", "description": "Résumé", "content": "Déroulé détaillé."}}
|
||||
```
|
||||
```loremind-action
|
||||
{{"type": "chapter", "name": "Nom", "description": "Résumé du chapitre."}}
|
||||
```
|
||||
```loremind-action
|
||||
{{"type": "arc", "name": "Nom", "description": "Résumé", "arcType": "LINEAR"}}
|
||||
```
|
||||
```loremind-action
|
||||
{{"type": "table", "name": "Nom", "diceFormula": "1d8", "entries": [{{"minRoll":1,"maxRoll":4,"label":"...","detail":"..."}}]}}
|
||||
```
|
||||
|
||||
Réponds en français, de façon utile et concise. Mets le texte explicatif AVANT les blocs d'action."""
|
||||
|
||||
|
||||
class NotebookChatUseCase:
|
||||
def __init__(self, rag: NotebookRagUseCase, llm: LLMChatProvider) -> None:
|
||||
self._rag = rag
|
||||
self._llm = llm
|
||||
|
||||
async def stream(
|
||||
self,
|
||||
source_ids: list[str],
|
||||
messages: list[ChatMessage],
|
||||
context: str = "",
|
||||
top_k: int = 6,
|
||||
) -> AsyncIterator[str]:
|
||||
last_user = next((m.content for m in reversed(messages) if m.role == "user"), "")
|
||||
passages = await self._rag.retrieve(source_ids, last_user, top_k=top_k)
|
||||
sources_block = (
|
||||
"\n\n".join(self._format_passage(p) for p in passages)
|
||||
if passages else "(aucun passage pertinent trouvé dans les sources)"
|
||||
)
|
||||
context_block = (
|
||||
f"--- TA CAMPAGNE ---\n{context.strip()}\n--- FIN CAMPAGNE ---\n\n"
|
||||
if context.strip() else "--- TA CAMPAGNE ---\n(aucune donnée de campagne)\n--- FIN CAMPAGNE ---\n\n"
|
||||
)
|
||||
system_prompt = _SYSTEM_PROMPT.format(
|
||||
context_block=context_block, sources_block=sources_block)
|
||||
async for token in self._llm.stream_chat(messages, system_prompt=system_prompt):
|
||||
yield token
|
||||
|
||||
@staticmethod
|
||||
def _format_passage(p: dict) -> str:
|
||||
page = p.get("page")
|
||||
prefix = f"(p. {page}) " if page else ""
|
||||
return f"• {prefix}{p['text'].strip()}"
|
||||
127
brain/app/application/notebook_deep.py
Normal file
127
brain/app/application/notebook_deep.py
Normal file
@@ -0,0 +1,127 @@
|
||||
"""Use case « Analyse approfondie » d'un notebook : map-reduce sur TOUT le document.
|
||||
|
||||
Contrairement au chat RAG (qui ne ramène que les top-k extraits), ce mode lit
|
||||
l'INTÉGRALITÉ des sources par lots :
|
||||
- MAP : pour chaque lot, le modèle extrait ce qui est pertinent pour la question
|
||||
(ou « RAS » si rien) ;
|
||||
- REDUCE : il synthétise toutes les notes en une réponse finale (streamée).
|
||||
|
||||
→ Répond aux questions globales/exhaustives (« liste tous les… ») quel que soit le
|
||||
modèle, au prix de plusieurs appels (comme l'import). Le lot est dimensionné par
|
||||
`batch_tokens` (= taille de morceau d'import) : avec un modèle gros-contexte, peu de
|
||||
lots ; avec un petit modèle local, plus de lots (mais ça reste exhaustif).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import AsyncIterator
|
||||
|
||||
import tiktoken
|
||||
|
||||
from app.application.llm_retry import generate_with_retry
|
||||
from app.domain.models import ChatMessage
|
||||
from app.domain.ports import LLMChatProvider, LLMProvider, LLMProviderError
|
||||
from app.infrastructure import vector_store
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_NO_MATCH = "RAS"
|
||||
_MAP_TEMPERATURE = 0.2
|
||||
|
||||
_MAP_PROMPT = """Voici un EXTRAIT d'un document. Extrais UNIQUEMENT les informations
|
||||
pertinentes pour répondre à la question ci-dessous. Conserve les détails utiles et
|
||||
indique les numéros de page (format « p. X »). Si l'extrait ne contient RIEN de
|
||||
pertinent, réponds EXACTEMENT « {no_match} » et rien d'autre.
|
||||
|
||||
QUESTION : {question}
|
||||
|
||||
--- EXTRAIT ---
|
||||
{excerpt}
|
||||
--- FIN EXTRAIT ---
|
||||
|
||||
Informations pertinentes (ou « {no_match} ») :"""
|
||||
|
||||
_REDUCE_SYSTEM = """Tu réponds à la question d'un MJ à partir de NOTES extraites de
|
||||
l'ENSEMBLE d'un document source (donc tu as une vue COMPLÈTE, pas un simple extrait).
|
||||
Synthétise ces notes en une réponse claire et structurée, cite les pages (« p. X »),
|
||||
et n'invente rien qui n'y figure pas. Si une CAMPAGNE est fournie ci-dessous, relie ta
|
||||
réponse à sa structure / ses PNJ pour des adaptations cohérentes.
|
||||
|
||||
{context_block}
|
||||
--- NOTES EXTRAITES DE TOUT LE DOCUMENT ---
|
||||
{notes_block}
|
||||
--- FIN DES NOTES ---
|
||||
|
||||
Réponds en français."""
|
||||
|
||||
|
||||
class NotebookDeepUseCase:
|
||||
def __init__(self, llm: LLMProvider, batch_tokens: int = 10000) -> None:
|
||||
self._llm = llm
|
||||
self._batch_tokens = max(2000, batch_tokens)
|
||||
|
||||
async def stream(
|
||||
self,
|
||||
source_ids: list[str],
|
||||
question: str,
|
||||
context: str = "",
|
||||
) -> AsyncIterator[dict]:
|
||||
"""Yield des évènements : {type:'progress',current,total}, {type:'token',token},
|
||||
{type:'done'}. (Les erreurs LLM des lots sont tolérées : lot ignoré.)"""
|
||||
chunks: list[dict] = []
|
||||
for sid in source_ids:
|
||||
chunks.extend(vector_store.all_chunks(sid))
|
||||
if not chunks:
|
||||
yield {"type": "token", "token": "Aucune source indexée à analyser."}
|
||||
yield {"type": "done"}
|
||||
return
|
||||
|
||||
batches = self._group(chunks)
|
||||
total = len(batches)
|
||||
notes: list[str] = []
|
||||
for i, batch in enumerate(batches):
|
||||
yield {"type": "progress", "current": i, "total": total}
|
||||
excerpt = "\n\n".join(
|
||||
f"(p. {c['page']}) {c['text'].strip()}" if c.get("page") else c["text"].strip()
|
||||
for c in batch
|
||||
)
|
||||
prompt = _MAP_PROMPT.format(no_match=_NO_MATCH, question=question, excerpt=excerpt)
|
||||
try:
|
||||
raw = await generate_with_retry(self._llm, prompt, temperature=_MAP_TEMPERATURE)
|
||||
except LLMProviderError as exc:
|
||||
logger.warning("Analyse approfondie : lot %s/%s ignoré : %s", i + 1, total, exc)
|
||||
continue
|
||||
answer = raw.strip()
|
||||
if answer and answer.upper().rstrip(".") != _NO_MATCH:
|
||||
notes.append(answer)
|
||||
yield {"type": "progress", "current": total, "total": total}
|
||||
|
||||
notes_block = "\n\n".join(notes) if notes else "(aucune information pertinente trouvée dans le document)"
|
||||
context_block = (
|
||||
f"--- TA CAMPAGNE (structure, PNJ, univers) ---\n{context.strip()}\n--- FIN CAMPAGNE ---\n\n"
|
||||
if context.strip() else ""
|
||||
)
|
||||
system_prompt = _REDUCE_SYSTEM.format(context_block=context_block, notes_block=notes_block)
|
||||
llm_chat: LLMChatProvider = self._llm # type: ignore[assignment]
|
||||
async for token in llm_chat.stream_chat(
|
||||
[ChatMessage(role="user", content=question)], system_prompt=system_prompt
|
||||
):
|
||||
yield {"type": "token", "token": token}
|
||||
yield {"type": "done"}
|
||||
|
||||
def _group(self, chunks: list[dict]) -> list[list[dict]]:
|
||||
"""Regroupe les extraits en lots ~`batch_tokens` (compte tiktoken)."""
|
||||
enc = tiktoken.get_encoding("cl100k_base")
|
||||
batches: list[list[dict]] = []
|
||||
current: list[dict] = []
|
||||
current_tokens = 0
|
||||
for c in chunks:
|
||||
t = len(enc.encode(c.get("text", "")))
|
||||
if current and current_tokens + t > self._batch_tokens:
|
||||
batches.append(current)
|
||||
current, current_tokens = [], 0
|
||||
current.append(c)
|
||||
current_tokens += t
|
||||
if current:
|
||||
batches.append(current)
|
||||
return batches
|
||||
79
brain/app/application/notebook_rag.py
Normal file
79
brain/app/application/notebook_rag.py
Normal file
@@ -0,0 +1,79 @@
|
||||
"""Use case RAG des notebooks : indexer une source PDF et retrouver les passages
|
||||
pertinents pour une question.
|
||||
|
||||
Chaîne d'indexation : PDF → extraction texte (+OCR) → découpage en extraits courts
|
||||
→ embeddings → stockage vectoriel (fichier). À la requête : on embed la question
|
||||
et on récupère les extraits les plus proches (cosinus) pour ancrer le chat.
|
||||
|
||||
Extraits PLUS COURTS que pour l'import (recopie) : ici on veut une granularité fine
|
||||
pour que la recherche pointe un passage précis, pas un demi-chapitre.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from app.application.chunking import chunk_text
|
||||
from app.application.embeddings import EmbeddingProvider
|
||||
from app.domain.ports import PdfTextExtractor
|
||||
from app.infrastructure import vector_store
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_RAG_CHUNK_TOKENS = 600
|
||||
# Un extrait avec quasi aucun texte réel (en-tête/pied de page, fragment de numéro
|
||||
# de page isolé « 249 250 ») ne sert à rien en RAG → on l'écarte. Seuil bas et
|
||||
# conservateur : on ne coupe QUE les fragments quasi-vides, jamais une vraie phrase.
|
||||
_MIN_LETTERS = 15
|
||||
|
||||
|
||||
def _has_enough_text(piece: str) -> bool:
|
||||
return sum(c.isalpha() for c in piece) >= _MIN_LETTERS
|
||||
|
||||
|
||||
class NotebookRagUseCase:
|
||||
def __init__(
|
||||
self,
|
||||
extractor: PdfTextExtractor,
|
||||
embedder: EmbeddingProvider,
|
||||
chunk_target_tokens: int = _RAG_CHUNK_TOKENS,
|
||||
) -> None:
|
||||
self._extractor = extractor
|
||||
self._embedder = embedder
|
||||
self._chunk_target_tokens = chunk_target_tokens
|
||||
|
||||
async def index_source(self, source_id: str, pdf_bytes: bytes) -> dict:
|
||||
"""Extrait, découpe PAR PAGE (pour garder le n° de page → citations), embed
|
||||
et stocke une source. Renvoie un récap."""
|
||||
doc = self._extractor.extract(pdf_bytes)
|
||||
chunks: list[str] = []
|
||||
pages: list[int] = []
|
||||
for page in doc.pages:
|
||||
for piece in chunk_text(page.text, self._chunk_target_tokens):
|
||||
if not _has_enough_text(piece):
|
||||
continue # fragment quasi-vide (en-tête/pied/numéro) → ignoré
|
||||
chunks.append(piece)
|
||||
pages.append(page.index + 1) # n° de page 1-based pour l'affichage
|
||||
logger.info(
|
||||
"Indexation notebook source=%s : %s page(s) (%s OCR), %s extrait(s).",
|
||||
source_id, doc.page_count, doc.ocr_page_count, len(chunks),
|
||||
)
|
||||
if not chunks:
|
||||
vector_store.save(source_id, [], [])
|
||||
return {"chunks": 0, "page_count": doc.page_count, "ocr_page_count": doc.ocr_page_count}
|
||||
vectors = await self._embedder.embed(chunks)
|
||||
count = vector_store.save(source_id, chunks, vectors, pages)
|
||||
return {
|
||||
"chunks": count,
|
||||
"page_count": doc.page_count,
|
||||
"ocr_page_count": doc.ocr_page_count,
|
||||
}
|
||||
|
||||
async def retrieve(self, source_ids: list[str], query: str, top_k: int = 6) -> list[dict]:
|
||||
"""Passages les plus pertinents (toutes sources) pour `query`."""
|
||||
ids = [s for s in source_ids if vector_store.exists(s)]
|
||||
if not ids or not query.strip():
|
||||
return []
|
||||
query_vectors = await self._embedder.embed([query])
|
||||
if not query_vectors:
|
||||
return []
|
||||
return vector_store.search(ids, query_vectors[0], top_k)
|
||||
45
brain/app/application/streaming.py
Normal file
45
brain/app/application/streaming.py
Normal file
@@ -0,0 +1,45 @@
|
||||
"""Heartbeats pour garder un flux SSE 'vivant' pendant une coroutine longue.
|
||||
|
||||
Problème résolu : pendant un appel LLM lent (import sur provider gratuit), le
|
||||
Brain ne produit AUCUN évènement SSE. Le Core (WebClient) ne 'voit aucun item'
|
||||
et coupe la connexion sur timeout d'inactivité :
|
||||
|
||||
ReactiveException: Did not observe any item or terminal signal within Nms
|
||||
|
||||
C'est le piège classique du SSE long. La parade standard = envoyer un keep-alive
|
||||
périodique. `with_heartbeat` exécute une coroutine en émettant un évènement
|
||||
'heartbeat' toutes les `interval` secondes tant qu'elle tourne, puis son résultat
|
||||
('result', valeur). Le Core remet son chrono à zéro sur n'importe quel évènement
|
||||
reçu (même inconnu) → plus de coupure, quelle que soit la lenteur du modèle.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from typing import Any, AsyncIterator, Awaitable
|
||||
|
||||
# Bien sous le timeout d'inactivité du Core (600s) ET de tout proxy (nginx ~60s).
|
||||
HEARTBEAT_INTERVAL_SECONDS = 15.0
|
||||
|
||||
|
||||
async def with_heartbeat(
|
||||
coro: Awaitable[Any],
|
||||
*,
|
||||
interval: float = HEARTBEAT_INTERVAL_SECONDS,
|
||||
) -> AsyncIterator[tuple[str, Any]]:
|
||||
"""Exécute `coro` en émettant ('heartbeat', None) toutes les `interval`s tant
|
||||
qu'elle n'est pas terminée, puis ('result', valeur).
|
||||
|
||||
L'exception éventuelle de `coro` est propagée (re-levée par `task.result()`),
|
||||
donc l'appelant peut l'attraper normalement. Si l'itération est abandonnée
|
||||
(client déconnecté), la tâche sous-jacente est annulée.
|
||||
"""
|
||||
task: asyncio.Task = asyncio.ensure_future(coro)
|
||||
try:
|
||||
while not task.done():
|
||||
done, _ = await asyncio.wait({task}, timeout=interval)
|
||||
if not done:
|
||||
yield ("heartbeat", None)
|
||||
yield ("result", task.result())
|
||||
finally:
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
@@ -25,8 +25,9 @@ class Settings(BaseSettings):
|
||||
extra="ignore",
|
||||
)
|
||||
|
||||
# Provider LLM actif. "ollama" = local ; "onemin" = 1min.ai ; "openrouter" = OpenRouter.
|
||||
llm_provider: Literal["ollama", "onemin", "openrouter"] = "ollama"
|
||||
# Provider LLM actif. "ollama" = local ; "onemin" = 1min.ai ;
|
||||
# "openrouter" = OpenRouter ; "mistral" = Mistral ; "gemini" = Google Gemini.
|
||||
llm_provider: Literal["ollama", "onemin", "openrouter", "mistral", "gemini"] = "ollama"
|
||||
|
||||
ollama_base_url: str = "http://localhost:11434"
|
||||
llm_model: str = "gemma4:26b"
|
||||
@@ -53,6 +54,35 @@ class Settings(BaseSettings):
|
||||
openrouter_api_key: str = ""
|
||||
openrouter_model: str = "openrouter/free"
|
||||
|
||||
# Mistral (La Plateforme, OpenAI-compatible). Cle + modele modifiables depuis
|
||||
# l'UI. Tier gratuit « Experiment » sur console.mistral.ai (sans CB). Defaut =
|
||||
# mistral-large-latest (128k contexte, bon en francais et en JSON fidele).
|
||||
mistral_api_key: str = ""
|
||||
mistral_model: str = "mistral-large-latest"
|
||||
|
||||
# Google Gemini (endpoint OpenAI-compatible). Cle gratuite sur
|
||||
# aistudio.google.com (sans CB). Defaut = gemini-2.0-flash : ~1M de contexte
|
||||
# (un livre tient en 1-2 appels), rapide, fidele, quota gratuit genereux.
|
||||
gemini_api_key: str = ""
|
||||
gemini_model: str = "gemini-2.0-flash"
|
||||
|
||||
# Embeddings (RAG des notebooks/ateliers). Modele SEPARE du chat.
|
||||
# "ollama" = local (gratuit, illimite, ideal pour indexer un livre = bcp
|
||||
# d'appels) ; "mistral" = cloud EU (mistral-embed, soumis au rate limit).
|
||||
embedding_provider: Literal["ollama", "mistral"] = "ollama"
|
||||
ollama_embedding_model: str = "nomic-embed-text"
|
||||
mistral_embedding_model: str = "mistral-embed"
|
||||
# Au démarrage, si le provider d'embeddings est Ollama et que le modèle n'est
|
||||
# pas présent, le Brain le télécharge automatiquement (en arrière-plan) → le RAG
|
||||
# marche "out of the box" pour un nouvel utilisateur. Désactivable (connexion
|
||||
# limitée, gestion manuelle des modèles).
|
||||
auto_pull_embedding_model: bool = True
|
||||
|
||||
# Nombre d'extraits récupérés par question dans le chat des ateliers (RAG).
|
||||
# Plus haut = plus de couverture pour les questions larges (« liste les… »),
|
||||
# mais prompt plus long. 8 par défaut (montable jusqu'à ~20 sur grand contexte).
|
||||
rag_top_k: int = 8
|
||||
|
||||
# Taille cible d'un morceau (en tokens) pour l'import de PDF (regles/campagne).
|
||||
# Plus c'est gros, moins il y a de morceaux => moins de fragmentation et un
|
||||
# import plus rapide, MAIS il faut que ca tienne dans la fenetre du modele.
|
||||
|
||||
@@ -31,6 +31,15 @@ _ALLOWED_KEYS = frozenset({
|
||||
"onemin_model",
|
||||
"openrouter_api_key",
|
||||
"openrouter_model",
|
||||
"mistral_api_key",
|
||||
"mistral_model",
|
||||
"gemini_api_key",
|
||||
"gemini_model",
|
||||
"embedding_provider",
|
||||
"ollama_embedding_model",
|
||||
"mistral_embedding_model",
|
||||
"auto_pull_embedding_model",
|
||||
"rag_top_k",
|
||||
"import_chunk_tokens",
|
||||
})
|
||||
|
||||
|
||||
178
brain/app/infrastructure/gemini_adapter.py
Normal file
178
brain/app/infrastructure/gemini_adapter.py
Normal file
@@ -0,0 +1,178 @@
|
||||
"""Adapter Google Gemini — implémente les ports LLMProvider / LLMChatProvider.
|
||||
|
||||
Gemini expose un endpoint COMPATIBLE OpenAI
|
||||
(POST {base}/openai/chat/completions, SSE), donc cet adapter est un client
|
||||
"OpenAI-compatible" — même structure que les adapters OpenRouter / Mistral.
|
||||
|
||||
Tier GRATUIT : clé API sur aistudio.google.com (sans CB). Atout majeur pour
|
||||
l'extraction de PDF : un CONTEXTE de ~1M tokens → un livre entier tient en 1-2
|
||||
appels, donc quasi aucun morceau perdu et peu de requêtes (limites jamais
|
||||
atteintes). Modèle conseillé : `gemini-2.0-flash` (rapide, gros contexte, fidèle).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from typing import AsyncIterator
|
||||
|
||||
import httpx
|
||||
|
||||
from app.core.config import Settings
|
||||
from app.domain.models import ChatMessage
|
||||
from app.domain.ports import LLMProviderError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_API_URL = "https://generativelanguage.googleapis.com/v1beta/openai/chat/completions"
|
||||
|
||||
# Délai max pour le PREMIER token de contenu (échec rapide si le modèle ne produit
|
||||
# rien). Gemini répond vite ; 120s est large.
|
||||
_FIRST_TOKEN_TIMEOUT_SECONDS = 120.0
|
||||
|
||||
|
||||
class GeminiLLMProvider:
|
||||
"""Adapter Gemini (OpenAI-compatible) — satisfait LLMProvider et LLMChatProvider."""
|
||||
|
||||
def __init__(self, settings: Settings) -> None:
|
||||
if not settings.gemini_api_key:
|
||||
raise LLMProviderError(
|
||||
"Clé API Gemini manquante. Configure-la depuis l'écran Paramètres "
|
||||
"(clé gratuite sur aistudio.google.com)."
|
||||
)
|
||||
self._api_key = settings.gemini_api_key
|
||||
self._model = settings.gemini_model
|
||||
self._timeout = settings.llm_timeout_seconds
|
||||
|
||||
def _headers(self) -> dict[str, str]:
|
||||
return {
|
||||
"Authorization": f"Bearer {self._api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
|
||||
async def generate(
|
||||
self,
|
||||
prompt: str,
|
||||
*,
|
||||
output_format: str | None = None,
|
||||
temperature: float | None = None,
|
||||
) -> str:
|
||||
"""One-shot via streaming (puis recollage), avec garde-fous au temps écoulé."""
|
||||
return await self._collect_with_timeouts(
|
||||
[ChatMessage(role="user", content=prompt)], temperature, output_format
|
||||
)
|
||||
|
||||
async def _collect_with_timeouts(
|
||||
self,
|
||||
messages: list[ChatMessage],
|
||||
temperature: float | None,
|
||||
output_format: str | None,
|
||||
) -> str:
|
||||
"""Collecte le stream avec deux garde-fous : 1er token borné (échec rapide
|
||||
si rien ne sort) + ceiling global `self._timeout`."""
|
||||
async def _collect() -> str:
|
||||
chunks: list[str] = []
|
||||
agen = self._stream(messages, None, temperature, output_format)
|
||||
try:
|
||||
while True:
|
||||
first = _FIRST_TOKEN_TIMEOUT_SECONDS if not chunks else None
|
||||
try:
|
||||
token = await asyncio.wait_for(agen.__anext__(), timeout=first)
|
||||
except StopAsyncIteration:
|
||||
break
|
||||
except asyncio.TimeoutError:
|
||||
raise LLMProviderError(
|
||||
f"Erreur Gemini : aucun contenu produit en "
|
||||
f"{int(_FIRST_TOKEN_TIMEOUT_SECONDS)}s. Réessayez ou vérifiez "
|
||||
"votre quota gratuit."
|
||||
)
|
||||
chunks.append(token)
|
||||
finally:
|
||||
await agen.aclose()
|
||||
return "".join(chunks)
|
||||
|
||||
try:
|
||||
return await asyncio.wait_for(_collect(), timeout=self._timeout)
|
||||
except asyncio.TimeoutError as exc:
|
||||
raise LLMProviderError(
|
||||
f"Erreur Gemini : génération non terminée en {self._timeout}s. Réduisez la "
|
||||
"taille des morceaux d'import ou augmentez le timeout."
|
||||
) from exc
|
||||
|
||||
async def stream_chat(
|
||||
self,
|
||||
messages: list[ChatMessage],
|
||||
*,
|
||||
system_prompt: str | None = None,
|
||||
temperature: float | None = None,
|
||||
) -> AsyncIterator[str]:
|
||||
async for token in self._stream(messages, system_prompt, temperature):
|
||||
yield token
|
||||
|
||||
async def _stream(
|
||||
self,
|
||||
messages: list[ChatMessage],
|
||||
system_prompt: str | None,
|
||||
temperature: float | None,
|
||||
output_format: str | None = None,
|
||||
) -> AsyncIterator[str]:
|
||||
payload_messages: list[dict[str, str]] = []
|
||||
if system_prompt:
|
||||
payload_messages.append({"role": "system", "content": system_prompt})
|
||||
for m in messages:
|
||||
payload_messages.append({"role": m.role, "content": m.content})
|
||||
|
||||
body: dict[str, object] = {
|
||||
"model": self._model,
|
||||
"messages": payload_messages,
|
||||
"stream": True,
|
||||
}
|
||||
if temperature is not None:
|
||||
body["temperature"] = temperature
|
||||
|
||||
async with httpx.AsyncClient(timeout=self._timeout) as client:
|
||||
try:
|
||||
async with client.stream(
|
||||
"POST", _API_URL, headers=self._headers(), json=body
|
||||
) as response:
|
||||
if response.status_code >= 400:
|
||||
detail = (await response.aread()).decode("utf-8", "replace").strip()
|
||||
raise LLMProviderError(
|
||||
f"Erreur Gemini (HTTP {response.status_code})"
|
||||
+ (f" : {detail[:500]}" if detail else "")
|
||||
)
|
||||
async for token in self._parse_sse(response):
|
||||
yield token
|
||||
except httpx.HTTPError as exc:
|
||||
raise LLMProviderError(self._format_http_error(exc)) from exc
|
||||
|
||||
@staticmethod
|
||||
async def _parse_sse(response: httpx.Response) -> AsyncIterator[str]:
|
||||
"""SSE OpenAI : lignes `data: {json}`, fin sur `data: [DONE]`."""
|
||||
async for line in response.aiter_lines():
|
||||
if not line or not line.startswith("data:"):
|
||||
continue
|
||||
data = line[len("data:"):].strip()
|
||||
if data == "[DONE]":
|
||||
return
|
||||
try:
|
||||
obj = json.loads(data)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
choices = obj.get("choices")
|
||||
if not choices:
|
||||
continue
|
||||
delta = choices[0].get("delta") or {}
|
||||
content = delta.get("content")
|
||||
if content:
|
||||
yield content
|
||||
|
||||
def _format_http_error(self, exc: httpx.HTTPError) -> str:
|
||||
if isinstance(exc, httpx.TimeoutException):
|
||||
return (
|
||||
f"Erreur Gemini : délai dépassé (timeout {self._timeout}s). Le modèle a "
|
||||
"mis trop de temps — réduis la taille des morceaux d'import ou augmente le timeout."
|
||||
)
|
||||
detail = str(exc) or exc.__class__.__name__
|
||||
return f"Erreur Gemini ({exc.__class__.__name__}) : {detail}"
|
||||
188
brain/app/infrastructure/mistral_adapter.py
Normal file
188
brain/app/infrastructure/mistral_adapter.py
Normal file
@@ -0,0 +1,188 @@
|
||||
"""Adapter Mistral — implémente les ports LLMProvider / LLMChatProvider.
|
||||
|
||||
Mistral (La Plateforme) expose l'API OpenAI standard (POST {base}/chat/completions,
|
||||
SSE), donc cet adapter est un client "OpenAI-compatible" — même structure que
|
||||
l'adapter OpenRouter. Le `generate` one-shot passe par le streaming (puis
|
||||
recollage) avec un timeout au temps écoulé pour ne jamais pendre à l'infini.
|
||||
|
||||
Tier GRATUIT : compte sur console.mistral.ai (tier « Experiment »), clé API à
|
||||
coller dans l'écran Paramètres. Modèles conseillés pour l'extraction : un grand
|
||||
contexte fidèle comme `mistral-large-latest` (128k) ou `mistral-small-latest`.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from typing import AsyncIterator
|
||||
|
||||
import httpx
|
||||
|
||||
from app.core.config import Settings
|
||||
from app.domain.models import ChatMessage
|
||||
from app.domain.ports import LLMProviderError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_API_URL = "https://api.mistral.ai/v1/chat/completions"
|
||||
|
||||
# Délai max pour le PREMIER token de contenu (échec rapide si le modèle est en file
|
||||
# d'attente et n'envoie que des keep-alive). Généreux car la file d'un tier gratuit
|
||||
# peut être longue.
|
||||
_FIRST_TOKEN_TIMEOUT_SECONDS = 120.0
|
||||
|
||||
|
||||
class MistralLLMProvider:
|
||||
"""Adapter Mistral (OpenAI-compatible) — satisfait LLMProvider et LLMChatProvider."""
|
||||
|
||||
def __init__(self, settings: Settings) -> None:
|
||||
if not settings.mistral_api_key:
|
||||
raise LLMProviderError(
|
||||
"Clé API Mistral manquante. Configure-la depuis l'écran Paramètres."
|
||||
)
|
||||
self._api_key = settings.mistral_api_key
|
||||
self._model = settings.mistral_model
|
||||
self._timeout = settings.llm_timeout_seconds
|
||||
|
||||
def _headers(self) -> dict[str, str]:
|
||||
return {
|
||||
"Authorization": f"Bearer {self._api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "application/json",
|
||||
}
|
||||
|
||||
async def generate(
|
||||
self,
|
||||
prompt: str,
|
||||
*,
|
||||
output_format: str | None = None,
|
||||
temperature: float | None = None,
|
||||
) -> str:
|
||||
"""One-shot via streaming (puis recollage) pour robustesse sur longues sorties.
|
||||
|
||||
Timeout au TEMPS ÉCOULÉ (asyncio) en plus du timeout réseau d'httpx :
|
||||
si le provider envoyait des keep-alive sans contenu, l'appel pendrait à
|
||||
l'infini. Ici on coupe net après `self._timeout` secondes.
|
||||
"""
|
||||
return await self._collect_with_timeouts(
|
||||
[ChatMessage(role="user", content=prompt)], temperature, output_format
|
||||
)
|
||||
|
||||
async def _collect_with_timeouts(
|
||||
self,
|
||||
messages: list[ChatMessage],
|
||||
temperature: float | None,
|
||||
output_format: str | None,
|
||||
) -> str:
|
||||
"""Collecte le stream avec deux garde-fous au temps écoulé : 1er token borné
|
||||
(file d'attente → échec rapide) + ceiling global `self._timeout`."""
|
||||
async def _collect() -> str:
|
||||
chunks: list[str] = []
|
||||
agen = self._stream(messages, None, temperature, output_format)
|
||||
try:
|
||||
while True:
|
||||
first = _FIRST_TOKEN_TIMEOUT_SECONDS if not chunks else None
|
||||
try:
|
||||
token = await asyncio.wait_for(agen.__anext__(), timeout=first)
|
||||
except StopAsyncIteration:
|
||||
break
|
||||
except asyncio.TimeoutError:
|
||||
raise LLMProviderError(
|
||||
f"Erreur Mistral : aucun contenu produit en "
|
||||
f"{int(_FIRST_TOKEN_TIMEOUT_SECONDS)}s — le modèle est probablement "
|
||||
"en file d'attente (tier gratuit, 2 req/min). Réessayez plus tard ou "
|
||||
"choisissez un modèle plus disponible."
|
||||
)
|
||||
chunks.append(token)
|
||||
finally:
|
||||
await agen.aclose()
|
||||
return "".join(chunks)
|
||||
|
||||
try:
|
||||
return await asyncio.wait_for(_collect(), timeout=self._timeout)
|
||||
except asyncio.TimeoutError as exc:
|
||||
raise LLMProviderError(
|
||||
f"Erreur Mistral : génération non terminée en {self._timeout}s. Réduisez la "
|
||||
"taille des morceaux d'import, augmentez le timeout, ou changez de modèle."
|
||||
) from exc
|
||||
|
||||
async def stream_chat(
|
||||
self,
|
||||
messages: list[ChatMessage],
|
||||
*,
|
||||
system_prompt: str | None = None,
|
||||
temperature: float | None = None,
|
||||
) -> AsyncIterator[str]:
|
||||
async for token in self._stream(messages, system_prompt, temperature):
|
||||
yield token
|
||||
|
||||
async def _stream(
|
||||
self,
|
||||
messages: list[ChatMessage],
|
||||
system_prompt: str | None,
|
||||
temperature: float | None,
|
||||
output_format: str | None = None,
|
||||
) -> AsyncIterator[str]:
|
||||
payload_messages: list[dict[str, str]] = []
|
||||
if system_prompt:
|
||||
payload_messages.append({"role": "system", "content": system_prompt})
|
||||
for m in messages:
|
||||
payload_messages.append({"role": m.role, "content": m.content})
|
||||
|
||||
body: dict[str, object] = {
|
||||
"model": self._model,
|
||||
"messages": payload_messages,
|
||||
"stream": True,
|
||||
}
|
||||
if temperature is not None:
|
||||
body["temperature"] = temperature
|
||||
|
||||
async with httpx.AsyncClient(timeout=self._timeout) as client:
|
||||
try:
|
||||
async with client.stream(
|
||||
"POST", _API_URL, headers=self._headers(), json=body
|
||||
) as response:
|
||||
if response.status_code >= 400:
|
||||
# En streaming le corps n'est pas lu automatiquement : on le
|
||||
# lit pour exposer le détail de Mistral (modèle inconnu, clé
|
||||
# invalide 401, quota 429…), sinon on n'a que le code HTTP nu.
|
||||
detail = (await response.aread()).decode("utf-8", "replace").strip()
|
||||
raise LLMProviderError(
|
||||
f"Erreur Mistral (HTTP {response.status_code})"
|
||||
+ (f" : {detail[:500]}" if detail else "")
|
||||
)
|
||||
async for token in self._parse_sse(response):
|
||||
yield token
|
||||
except httpx.HTTPError as exc:
|
||||
raise LLMProviderError(self._format_http_error(exc)) from exc
|
||||
|
||||
@staticmethod
|
||||
async def _parse_sse(response: httpx.Response) -> AsyncIterator[str]:
|
||||
"""SSE OpenAI : lignes `data: {json}`, fin sur `data: [DONE]`."""
|
||||
async for line in response.aiter_lines():
|
||||
if not line or not line.startswith("data:"):
|
||||
continue # lignes vides ou keep-alive (`: ...`)
|
||||
data = line[len("data:"):].strip()
|
||||
if data == "[DONE]":
|
||||
return
|
||||
try:
|
||||
obj = json.loads(data)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
choices = obj.get("choices")
|
||||
if not choices:
|
||||
continue
|
||||
delta = choices[0].get("delta") or {}
|
||||
content = delta.get("content")
|
||||
if content:
|
||||
yield content
|
||||
|
||||
def _format_http_error(self, exc: httpx.HTTPError) -> str:
|
||||
"""Message lisible (timeout, quota 429, clé invalide 401, modèle inconnu…)."""
|
||||
if isinstance(exc, httpx.TimeoutException):
|
||||
return (
|
||||
f"Erreur Mistral : délai dépassé (timeout {self._timeout}s). Le modèle a "
|
||||
"mis trop de temps — réduis la taille des morceaux d'import ou augmente le timeout."
|
||||
)
|
||||
detail = str(exc) or exc.__class__.__name__
|
||||
return f"Erreur Mistral ({exc.__class__.__name__}) : {detail}"
|
||||
58
brain/app/infrastructure/mistral_embedding_adapter.py
Normal file
58
brain/app/infrastructure/mistral_embedding_adapter.py
Normal file
@@ -0,0 +1,58 @@
|
||||
"""Adapter d'embeddings Mistral (cloud, EU) — POST /v1/embeddings.
|
||||
|
||||
Soumis au rate limit du tier gratuit : pour indexer un gros document on envoie
|
||||
les textes par lots (et l'appelant peut espacer les appels si besoin).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import httpx
|
||||
|
||||
from app.application.embeddings import EmbeddingError
|
||||
from app.core.config import Settings
|
||||
|
||||
_API_URL = "https://api.mistral.ai/v1/embeddings"
|
||||
# Lot raisonnable pour ne pas envoyer un payload géant d'un coup.
|
||||
_BATCH_SIZE = 64
|
||||
|
||||
|
||||
class MistralEmbeddingProvider:
|
||||
"""Implémente EmbeddingProvider via l'API Mistral embeddings."""
|
||||
|
||||
def __init__(self, settings: Settings) -> None:
|
||||
if not settings.mistral_api_key:
|
||||
raise EmbeddingError(
|
||||
"Clé API Mistral manquante (requise pour les embeddings Mistral). "
|
||||
"Configure-la dans les Paramètres ou choisis Ollama pour les embeddings."
|
||||
)
|
||||
self._api_key = settings.mistral_api_key
|
||||
self._model = settings.mistral_embedding_model
|
||||
self._timeout = settings.llm_timeout_seconds
|
||||
|
||||
async def embed(self, texts: list[str]) -> list[list[float]]:
|
||||
if not texts:
|
||||
return []
|
||||
out: list[list[float]] = []
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self._api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
async with httpx.AsyncClient(timeout=self._timeout) as client:
|
||||
for start in range(0, len(texts), _BATCH_SIZE):
|
||||
batch = texts[start:start + _BATCH_SIZE]
|
||||
try:
|
||||
response = await client.post(
|
||||
_API_URL, headers=headers, json={"model": self._model, "input": batch})
|
||||
if response.status_code >= 400:
|
||||
raise EmbeddingError(
|
||||
f"Mistral embeddings HTTP {response.status_code} : "
|
||||
f"{response.text.strip()[:300]}")
|
||||
data = response.json()
|
||||
except httpx.HTTPError as exc:
|
||||
raise EmbeddingError(f"Erreur Mistral embeddings : {exc}") from exc
|
||||
|
||||
items = data.get("data")
|
||||
if not isinstance(items, list) or len(items) != len(batch):
|
||||
raise EmbeddingError("Réponse d'embeddings Mistral inattendue (taille incohérente).")
|
||||
for item in items:
|
||||
out.append([float(x) for x in item.get("embedding", [])])
|
||||
return out
|
||||
43
brain/app/infrastructure/ollama_embedding_adapter.py
Normal file
43
brain/app/infrastructure/ollama_embedding_adapter.py
Normal file
@@ -0,0 +1,43 @@
|
||||
"""Adapter d'embeddings Ollama (local) — endpoint /api/embed.
|
||||
|
||||
Gratuit et illimité (tourne sur la machine). Nécessite d'avoir pullé le modèle
|
||||
d'embedding (ex. `ollama pull nomic-embed-text`).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import httpx
|
||||
|
||||
from app.application.embeddings import EmbeddingError
|
||||
from app.core.config import Settings
|
||||
|
||||
|
||||
class OllamaEmbeddingProvider:
|
||||
"""Implémente EmbeddingProvider via Ollama /api/embed (batch)."""
|
||||
|
||||
def __init__(self, settings: Settings) -> None:
|
||||
self._base_url = settings.ollama_base_url
|
||||
self._model = settings.ollama_embedding_model
|
||||
self._timeout = settings.llm_timeout_seconds
|
||||
|
||||
async def embed(self, texts: list[str]) -> list[list[float]]:
|
||||
if not texts:
|
||||
return []
|
||||
url = f"{self._base_url}/api/embed"
|
||||
payload = {"model": self._model, "input": texts}
|
||||
async with httpx.AsyncClient(timeout=self._timeout) as client:
|
||||
try:
|
||||
response = await client.post(url, json=payload)
|
||||
if response.status_code >= 400:
|
||||
body = response.text
|
||||
raise EmbeddingError(
|
||||
f"Ollama embeddings HTTP {response.status_code} : {body.strip()[:300]}. "
|
||||
f"Le modèle '{self._model}' est-il installé ? (ollama pull {self._model})"
|
||||
)
|
||||
data = response.json()
|
||||
except httpx.HTTPError as exc:
|
||||
raise EmbeddingError(f"Erreur Ollama embeddings : {exc}") from exc
|
||||
|
||||
vectors = data.get("embeddings")
|
||||
if not isinstance(vectors, list) or len(vectors) != len(texts):
|
||||
raise EmbeddingError("Réponse d'embeddings Ollama inattendue (taille incohérente).")
|
||||
return [[float(x) for x in v] for v in vectors]
|
||||
@@ -11,17 +11,26 @@ qui choisit automatiquement un modèle gratuit — aucun crédit consommé.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from typing import AsyncIterator
|
||||
|
||||
import httpx
|
||||
|
||||
from app.core.config import Settings
|
||||
from app.domain.models import ChatMessage
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
from app.domain.ports import LLMProviderError
|
||||
|
||||
_API_URL = "https://openrouter.ai/api/v1/chat/completions"
|
||||
|
||||
# Délai max pour le PREMIER token de contenu. Un modèle gratuit "en file d'attente"
|
||||
# n'envoie que des keep-alive (aucun contenu) → on échoue vite et clairement au lieu
|
||||
# de pendre. Généreux (2 min) car la file d'attente d'un tier gratuit peut être longue.
|
||||
_FIRST_TOKEN_TIMEOUT_SECONDS = 120.0
|
||||
|
||||
|
||||
class OpenRouterLLMProvider:
|
||||
"""Adapter OpenRouter (OpenAI-compatible) — satisfait LLMProvider et LLMChatProvider."""
|
||||
@@ -51,11 +60,63 @@ class OpenRouterLLMProvider:
|
||||
output_format: str | None = None,
|
||||
temperature: float | None = None,
|
||||
) -> str:
|
||||
"""One-shot via streaming (puis recollage) pour robustesse sur longues sorties."""
|
||||
chunks: list[str] = []
|
||||
async for token in self._stream([ChatMessage(role="user", content=prompt)], None, temperature):
|
||||
chunks.append(token)
|
||||
return "".join(chunks)
|
||||
"""One-shot via streaming (puis recollage) pour robustesse sur longues sorties.
|
||||
|
||||
Timeout au TEMPS ÉCOULÉ (asyncio) en plus du timeout réseau d'httpx : un
|
||||
modèle gratuit saturé/en file d'attente envoie des keep-alive (`: OPENROUTER
|
||||
PROCESSING`) mais AUCUN contenu → httpx ne déclenche jamais son read-timeout
|
||||
(des octets arrivent) et l'appel pendrait à l'infini. Ici on coupe net après
|
||||
`self._timeout` secondes, quoi qu'il arrive.
|
||||
"""
|
||||
return await self._collect_with_timeouts(
|
||||
[ChatMessage(role="user", content=prompt)], temperature, output_format, "OpenRouter"
|
||||
)
|
||||
|
||||
async def _collect_with_timeouts(
|
||||
self,
|
||||
messages: list[ChatMessage],
|
||||
temperature: float | None,
|
||||
output_format: str | None,
|
||||
provider: str,
|
||||
) -> str:
|
||||
"""Collecte le stream avec DEUX garde-fous au temps écoulé :
|
||||
- 1er token borné (`_FIRST_TOKEN_TIMEOUT_SECONDS`) : détecte un modèle bloqué
|
||||
en file d'attente (que des keep-alive, aucun contenu) → échec rapide ;
|
||||
- ceiling global (`self._timeout`) : génération qui ne se termine jamais.
|
||||
Le timeout réseau d'httpx ne suffit pas : des keep-alive font 'arriver des
|
||||
octets' et empêchent son read-timeout de se déclencher.
|
||||
"""
|
||||
async def _collect() -> str:
|
||||
chunks: list[str] = []
|
||||
agen = self._stream(messages, None, temperature, output_format)
|
||||
try:
|
||||
while True:
|
||||
# Borne SEULEMENT l'attente du 1er token (file d'attente) ; ensuite
|
||||
# on laisse générer (le ceiling global couvre le reste).
|
||||
first = _FIRST_TOKEN_TIMEOUT_SECONDS if not chunks else None
|
||||
try:
|
||||
token = await asyncio.wait_for(agen.__anext__(), timeout=first)
|
||||
except StopAsyncIteration:
|
||||
break
|
||||
except asyncio.TimeoutError:
|
||||
raise LLMProviderError(
|
||||
f"Erreur {provider} : aucun contenu produit en "
|
||||
f"{int(_FIRST_TOKEN_TIMEOUT_SECONDS)}s — le modèle gratuit est "
|
||||
"probablement en file d'attente / saturé. Réessayez plus tard ou "
|
||||
"choisissez un autre modèle (1min.ai, ou payant)."
|
||||
)
|
||||
chunks.append(token)
|
||||
finally:
|
||||
await agen.aclose()
|
||||
return "".join(chunks)
|
||||
|
||||
try:
|
||||
return await asyncio.wait_for(_collect(), timeout=self._timeout)
|
||||
except asyncio.TimeoutError as exc:
|
||||
raise LLMProviderError(
|
||||
f"Erreur {provider} : génération non terminée en {self._timeout}s. Réduisez la "
|
||||
"taille des morceaux d'import, augmentez le timeout, ou changez de modèle."
|
||||
) from exc
|
||||
|
||||
async def stream_chat(
|
||||
self,
|
||||
@@ -72,6 +133,7 @@ class OpenRouterLLMProvider:
|
||||
messages: list[ChatMessage],
|
||||
system_prompt: str | None,
|
||||
temperature: float | None,
|
||||
output_format: str | None = None,
|
||||
) -> AsyncIterator[str]:
|
||||
payload_messages: list[dict[str, str]] = []
|
||||
if system_prompt:
|
||||
@@ -86,6 +148,10 @@ class OpenRouterLLMProvider:
|
||||
}
|
||||
if temperature is not None:
|
||||
body["temperature"] = temperature
|
||||
# NB : on n'impose PAS `response_format=json_object`. Beaucoup de modèles/
|
||||
# providers GRATUITS ne le supportent pas et renvoient une réponse VIDE.
|
||||
# On laisse le modèle répondre librement ; l'extraction JSON en aval
|
||||
# (load_json_object + nettoyage du raisonnement) récupère le JSON dans la prose.
|
||||
|
||||
async with httpx.AsyncClient(timeout=self._timeout) as client:
|
||||
try:
|
||||
|
||||
109
brain/app/infrastructure/vector_store.py
Normal file
109
brain/app/infrastructure/vector_store.py
Normal file
@@ -0,0 +1,109 @@
|
||||
"""Stockage vectoriel fichier (RAG des notebooks) — sans dépendance lourde.
|
||||
|
||||
Chaque SOURCE est persistée en un fichier JSON sur le volume `data/` du Brain :
|
||||
data/notebooks/{source_id}.json = {"dim": N, "chunks": [{"text":..., "vector":[...]}]}
|
||||
|
||||
À l'échelle d'un livre (quelques centaines d'extraits), une recherche cosinus en
|
||||
Python pur est instantanée — inutile d'ajouter numpy/pgvector/une base vectorielle.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import math
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
_STORE_DIR = Path("data/notebooks")
|
||||
_SAFE_ID = re.compile(r"[^A-Za-z0-9_-]")
|
||||
|
||||
|
||||
def _path(source_id: str) -> Path:
|
||||
safe = _SAFE_ID.sub("_", str(source_id))
|
||||
return _STORE_DIR / f"{safe}.json"
|
||||
|
||||
|
||||
def save(
|
||||
source_id: str,
|
||||
chunks: list[str],
|
||||
vectors: list[list[float]],
|
||||
pages: list[int] | None = None,
|
||||
) -> int:
|
||||
"""Persiste les (chunk, vecteur[, page]) d'une source. Renvoie le nb d'extraits."""
|
||||
if len(chunks) != len(vectors):
|
||||
raise ValueError("chunks et vectors de tailles différentes")
|
||||
if pages is not None and len(pages) != len(chunks):
|
||||
raise ValueError("pages et chunks de tailles différentes")
|
||||
_STORE_DIR.mkdir(parents=True, exist_ok=True)
|
||||
items = []
|
||||
for i, (c, v) in enumerate(zip(chunks, vectors)):
|
||||
item = {"text": c, "vector": v}
|
||||
if pages is not None:
|
||||
item["page"] = pages[i]
|
||||
items.append(item)
|
||||
payload = {"dim": len(vectors[0]) if vectors else 0, "chunks": items}
|
||||
_path(source_id).write_text(json.dumps(payload, ensure_ascii=False), encoding="utf-8")
|
||||
return len(chunks)
|
||||
|
||||
|
||||
def exists(source_id: str) -> bool:
|
||||
return _path(source_id).exists()
|
||||
|
||||
|
||||
def delete(source_id: str) -> None:
|
||||
_path(source_id).unlink(missing_ok=True)
|
||||
|
||||
|
||||
def _load(source_id: str) -> list[dict]:
|
||||
p = _path(source_id)
|
||||
if not p.exists():
|
||||
return []
|
||||
try:
|
||||
data = json.loads(p.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError):
|
||||
return []
|
||||
return data.get("chunks", []) if isinstance(data, dict) else []
|
||||
|
||||
|
||||
def all_chunks(source_id: str) -> list[dict]:
|
||||
"""Tous les extraits d'une source (texte + page), sans vecteurs — pour le mode
|
||||
« analyse approfondie » (map-reduce sur tout le document)."""
|
||||
return [{"text": c.get("text", ""), "page": c.get("page")} for c in _load(source_id)]
|
||||
|
||||
|
||||
def _cosine(a: list[float], b: list[float]) -> float:
|
||||
if not a or not b or len(a) != len(b):
|
||||
return 0.0
|
||||
dot = 0.0
|
||||
na = 0.0
|
||||
nb = 0.0
|
||||
for x, y in zip(a, b):
|
||||
dot += x * y
|
||||
na += x * x
|
||||
nb += y * y
|
||||
if na == 0.0 or nb == 0.0:
|
||||
return 0.0
|
||||
return dot / (math.sqrt(na) * math.sqrt(nb))
|
||||
|
||||
|
||||
def search(
|
||||
source_ids: list[str],
|
||||
query_vector: list[float],
|
||||
top_k: int = 6,
|
||||
) -> list[dict]:
|
||||
"""Renvoie les `top_k` extraits les plus proches, toutes sources confondues.
|
||||
|
||||
Chaque résultat : {"text": str, "score": float, "source_id": str}.
|
||||
"""
|
||||
scored: list[dict] = []
|
||||
for sid in source_ids:
|
||||
for chunk in _load(sid):
|
||||
vector = chunk.get("vector") or []
|
||||
score = _cosine(query_vector, vector)
|
||||
scored.append({
|
||||
"text": chunk.get("text", ""),
|
||||
"score": score,
|
||||
"source_id": sid,
|
||||
"page": chunk.get("page"),
|
||||
})
|
||||
scored.sort(key=lambda c: c["score"], reverse=True)
|
||||
return scored[:top_k]
|
||||
@@ -4,7 +4,9 @@ Controller volontairement FIN : il valide l'entrée (DTOs Pydantic), délègue
|
||||
au domaine via injection de dépendance (ports + use cases), et transforme les
|
||||
erreurs du domaine en réponses HTTP. Aucune connaissance d'Ollama ici.
|
||||
"""
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from typing import Annotated, AsyncIterator, Literal
|
||||
|
||||
import hmac
|
||||
@@ -14,11 +16,22 @@ from fastapi import Depends, FastAPI, File, Form, HTTPException, Request, Upload
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
import re
|
||||
|
||||
from app.application.adapt_campaign import AdaptCampaignUseCase
|
||||
from app.application.chat import ChatUseCase
|
||||
from app.application.generate_page import GeneratePageUseCase
|
||||
from app.application.import_campaign import ImportCampaignUseCase
|
||||
from app.application.import_rules import ImportRulesUseCase
|
||||
from app.application.llm_json import load_json_object
|
||||
from app.application.llm_retry import generate_with_retry
|
||||
from app.application.notebook_rag import NotebookRagUseCase
|
||||
from app.application.notebook_chat import NotebookChatUseCase
|
||||
from app.application.notebook_deep import NotebookDeepUseCase
|
||||
from app.application.embeddings import EmbeddingError
|
||||
from app.infrastructure import vector_store
|
||||
from app.infrastructure.ollama_embedding_adapter import OllamaEmbeddingProvider
|
||||
from app.infrastructure.mistral_embedding_adapter import MistralEmbeddingProvider
|
||||
from app.core.config import Settings, get_settings
|
||||
from app.core.settings_store import save_overrides
|
||||
from app.domain.models import (
|
||||
@@ -46,14 +59,18 @@ from app.domain.ports import LLMProvider, LLMProviderError, PdfExtractionError
|
||||
from app.infrastructure.ollama_adapter import OllamaLLMProvider
|
||||
from app.infrastructure.onemin_adapter import OneMinAiLLMProvider
|
||||
from app.infrastructure.openrouter_adapter import OpenRouterLLMProvider
|
||||
from app.infrastructure.mistral_adapter import MistralLLMProvider
|
||||
from app.infrastructure.gemini_adapter import GeminiLLMProvider
|
||||
from app.infrastructure.pdf_extractor import PyMuPdfTextExtractor
|
||||
|
||||
app = FastAPI(
|
||||
title="LoreMind Brain",
|
||||
description="Backend IA pour la génération de contenu narratif.",
|
||||
version="0.10.3-beta",
|
||||
version="0.11.0-beta",
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# Encodeur tiktoken partagé — chargé une fois pour éviter le coût de lookup
|
||||
# à chaque requête. On utilise cl100k_base (GPT-3.5/4) comme tokenizer
|
||||
@@ -357,6 +374,10 @@ def get_llm_provider(
|
||||
return OneMinAiLLMProvider(settings)
|
||||
if settings.llm_provider == "openrouter":
|
||||
return OpenRouterLLMProvider(settings)
|
||||
if settings.llm_provider == "mistral":
|
||||
return MistralLLMProvider(settings)
|
||||
if settings.llm_provider == "gemini":
|
||||
return GeminiLLMProvider(settings)
|
||||
return OllamaLLMProvider(settings)
|
||||
except LLMProviderError as exc:
|
||||
# Ex : cle 1min.ai manquante. On renvoie du 400 plutot que du 500
|
||||
@@ -416,6 +437,38 @@ def get_adapt_campaign_use_case(
|
||||
llm=llm, extractor=_PDF_EXTRACTOR, max_input_tokens=settings.import_chunk_tokens)
|
||||
|
||||
|
||||
def get_embedding_provider(
|
||||
settings: Annotated[Settings, Depends(get_settings)],
|
||||
):
|
||||
"""Factory de l'adapter d'embeddings (RAG) selon `embedding_provider`."""
|
||||
try:
|
||||
if settings.embedding_provider == "mistral":
|
||||
return MistralEmbeddingProvider(settings)
|
||||
return OllamaEmbeddingProvider(settings)
|
||||
except EmbeddingError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
|
||||
def get_notebook_rag_use_case(
|
||||
embedder: Annotated[object, Depends(get_embedding_provider)],
|
||||
) -> NotebookRagUseCase:
|
||||
return NotebookRagUseCase(extractor=_PDF_EXTRACTOR, embedder=embedder) # type: ignore[arg-type]
|
||||
|
||||
|
||||
def get_notebook_chat_use_case(
|
||||
llm: Annotated[LLMProvider, Depends(get_llm_provider)],
|
||||
rag: Annotated[NotebookRagUseCase, Depends(get_notebook_rag_use_case)],
|
||||
) -> NotebookChatUseCase:
|
||||
return NotebookChatUseCase(rag=rag, llm=llm) # type: ignore[arg-type]
|
||||
|
||||
|
||||
def get_notebook_deep_use_case(
|
||||
llm: Annotated[LLMProvider, Depends(get_llm_provider)],
|
||||
settings: Annotated[Settings, Depends(get_settings)],
|
||||
) -> NotebookDeepUseCase:
|
||||
return NotebookDeepUseCase(llm=llm, batch_tokens=settings.import_chunk_tokens)
|
||||
|
||||
|
||||
# --- Endpoints ---
|
||||
|
||||
|
||||
@@ -425,6 +478,54 @@ def health() -> dict[str, str]:
|
||||
return {"status": "ok", "service": "brain"}
|
||||
|
||||
|
||||
@app.on_event("startup")
|
||||
async def _auto_install_embedding_model() -> None:
|
||||
"""Au démarrage : si le provider d'embeddings est Ollama et que le modèle n'est
|
||||
pas installé, on le télécharge EN ARRIÈRE-PLAN → le RAG marche d'emblée pour un
|
||||
nouvel utilisateur, sans bloquer le démarrage du Brain. Best-effort (Ollama peut
|
||||
être absent / la connexion limitée) ; désactivable via `auto_pull_embedding_model`.
|
||||
"""
|
||||
settings = get_settings()
|
||||
if not settings.auto_pull_embedding_model or settings.embedding_provider != "ollama":
|
||||
return
|
||||
asyncio.create_task(_ensure_ollama_embedding_model(settings.ollama_base_url, settings.ollama_embedding_model))
|
||||
|
||||
|
||||
async def _ensure_ollama_embedding_model(base_url: str, model: str) -> None:
|
||||
# Attend qu'Ollama soit joignable (ordre de démarrage des conteneurs), puis
|
||||
# vérifie la présence du modèle avant de le tirer.
|
||||
for attempt in range(10):
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=10) as client:
|
||||
tags = await client.get(f"{base_url}/api/tags")
|
||||
tags.raise_for_status()
|
||||
names = [m.get("name", "") for m in tags.json().get("models", [])]
|
||||
if any(n == model or n.startswith(model + ":") for n in names):
|
||||
logger.info("Modèle d'embedding '%s' déjà présent.", model)
|
||||
return
|
||||
break # Ollama joignable, modèle absent → on tire (ci-dessous)
|
||||
except httpx.HTTPError:
|
||||
await asyncio.sleep(min(5 * (attempt + 1), 30))
|
||||
else:
|
||||
logger.warning(
|
||||
"Ollama injoignable au démarrage — modèle d'embedding '%s' non auto-installé "
|
||||
"(il sera tirable manuellement : ollama pull %s).", model, model)
|
||||
return
|
||||
|
||||
logger.info("Téléchargement automatique du modèle d'embedding '%s'…", model)
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=None) as client:
|
||||
async with client.stream("POST", f"{base_url}/api/pull", json={"name": model}) as resp:
|
||||
resp.raise_for_status()
|
||||
async for _line in resp.aiter_lines():
|
||||
pass # on draine la progression NDJSON jusqu'à la fin
|
||||
logger.info("Modèle d'embedding '%s' prêt.", model)
|
||||
except httpx.HTTPError as exc:
|
||||
logger.warning(
|
||||
"Auto-installation du modèle d'embedding '%s' échouée : %s "
|
||||
"(tirage manuel possible : ollama pull %s).", model, exc, model)
|
||||
|
||||
|
||||
@app.post("/generate", response_model=GenerateResponse)
|
||||
async def generate(
|
||||
body: GenerateRequest,
|
||||
@@ -555,6 +656,11 @@ async def import_rules_stream(
|
||||
yield _sse("error", {"message": str(exc)})
|
||||
except LLMProviderError as exc:
|
||||
yield _sse("error", {"message": str(exc)})
|
||||
except Exception as exc: # noqa: BLE001 — filet : une erreur inattendue ne doit
|
||||
# PAS casser le flux SSE brutalement (sinon le Core n'a qu'un message générique
|
||||
# sans détail). On la transforme en évènement `error` propre + log avec trace.
|
||||
logger.exception("Import règles : erreur inattendue dans le flux.")
|
||||
yield _sse("error", {"message": f"Erreur inattendue du Brain : {type(exc).__name__} : {exc}"})
|
||||
|
||||
return StreamingResponse(event_stream(), media_type="text/event-stream")
|
||||
|
||||
@@ -590,6 +696,10 @@ async def import_campaign_stream(
|
||||
yield _sse("error", {"message": str(exc)})
|
||||
except LLMProviderError as exc:
|
||||
yield _sse("error", {"message": str(exc)})
|
||||
except Exception as exc: # noqa: BLE001 — voir import règles : on ne laisse pas
|
||||
# une erreur inattendue casser le flux sans détail.
|
||||
logger.exception("Import campagne : erreur inattendue dans le flux.")
|
||||
yield _sse("error", {"message": f"Erreur inattendue du Brain : {type(exc).__name__} : {exc}"})
|
||||
|
||||
return StreamingResponse(event_stream(), media_type="text/event-stream")
|
||||
|
||||
@@ -773,6 +883,245 @@ async def summarize_conversation_title(
|
||||
return SummarizeTitleResponseDTO(title=title)
|
||||
|
||||
|
||||
# --- Tables aléatoires : génération IA + improvisation -----------------------
|
||||
|
||||
_DICE_FORMULA_RE = re.compile(r"^\s*(\d*)\s*[dD]\s*(\d+)\s*$")
|
||||
|
||||
|
||||
def _dice_total_range(formula: str) -> tuple[int, int] | None:
|
||||
"""(min, max) des totaux possibles d'une formule NdM, ou None si invalide."""
|
||||
match = _DICE_FORMULA_RE.match(formula or "")
|
||||
if not match:
|
||||
return None
|
||||
count = int(match.group(1)) if match.group(1) else 1
|
||||
faces = int(match.group(2))
|
||||
if count < 1 or count > 100 or faces < 2 or faces > 10000:
|
||||
return None
|
||||
return count, count * faces
|
||||
|
||||
|
||||
class GenerateTableRequestDTO(BaseModel):
|
||||
description: str
|
||||
dice_formula: str = Field(default="1d20")
|
||||
# Contexte libre assemblé par le Core (nom de campagne, système, ambiance…).
|
||||
context: str = Field(default="")
|
||||
|
||||
|
||||
class GeneratedTableEntryDTO(BaseModel):
|
||||
min_roll: int
|
||||
max_roll: int
|
||||
label: str
|
||||
detail: str = ""
|
||||
|
||||
|
||||
class GenerateTableResponseDTO(BaseModel):
|
||||
name: str
|
||||
description: str = ""
|
||||
entries: list[GeneratedTableEntryDTO]
|
||||
|
||||
|
||||
@app.post("/generate/random-table", response_model=GenerateTableResponseDTO)
|
||||
async def generate_random_table(
|
||||
body: GenerateTableRequestDTO,
|
||||
llm: Annotated[LLMProvider, Depends(get_llm_provider)],
|
||||
) -> GenerateTableResponseDTO:
|
||||
"""Génère une table aléatoire (entrées par plage) couvrant la formule de dé."""
|
||||
rng = _dice_total_range(body.dice_formula)
|
||||
if rng is None:
|
||||
raise HTTPException(status_code=422, detail="Formule de dé invalide (ex. 1d20, 2d6, d100).")
|
||||
lo, hi = rng
|
||||
context_block = f"\nContexte de la campagne :\n{body.context.strip()}\n" if body.context.strip() else ""
|
||||
prompt = (
|
||||
"Tu es un assistant de jeu de rôle. Génère une TABLE ALÉATOIRE évocatrice.\n"
|
||||
f"Dé : {body.dice_formula} (résultats possibles de {lo} à {hi}).\n"
|
||||
f"Sujet : {body.description.strip()}\n"
|
||||
f"{context_block}\n"
|
||||
"Règles IMPÉRATIVES :\n"
|
||||
"- Réponds UNIQUEMENT par un objet JSON valide, sans texte autour.\n"
|
||||
'- Format : {"name": "...", "description": "...", "entries": '
|
||||
'[{"min_roll": N, "max_roll": M, "label": "résultat court", "detail": "1-2 phrases"}]}\n'
|
||||
f"- Les plages (min_roll..max_roll) doivent COUVRIR EXACTEMENT {lo}..{hi}, "
|
||||
"sans trou ni chevauchement, dans l'ordre croissant.\n"
|
||||
"- Des résultats variés, cohérents avec le sujet (et le contexte s'il est fourni).\n"
|
||||
"- En français. 'label' = résultat bref ; 'detail' = description/effet concret.\n"
|
||||
"Renvoie maintenant le JSON."
|
||||
)
|
||||
try:
|
||||
raw = await generate_with_retry(llm, prompt, output_format="json", temperature=0.7)
|
||||
except LLMProviderError as exc:
|
||||
raise HTTPException(status_code=502, detail=str(exc)) from exc
|
||||
|
||||
parsed, _ = load_json_object(raw)
|
||||
if not isinstance(parsed, dict):
|
||||
raise HTTPException(status_code=502, detail="Le modèle n'a pas renvoyé de table exploitable.")
|
||||
|
||||
entries: list[GeneratedTableEntryDTO] = []
|
||||
for e in parsed.get("entries", []) or []:
|
||||
if not isinstance(e, dict):
|
||||
continue
|
||||
try:
|
||||
mn = int(e["min_roll"])
|
||||
mx = int(e["max_roll"])
|
||||
except (KeyError, TypeError, ValueError):
|
||||
continue
|
||||
label = str(e.get("label") or "").strip()
|
||||
if not label:
|
||||
continue
|
||||
entries.append(GeneratedTableEntryDTO(
|
||||
min_roll=mn, max_roll=max(mn, mx), label=label[:200],
|
||||
detail=str(e.get("detail") or "").strip(),
|
||||
))
|
||||
if not entries:
|
||||
raise HTTPException(status_code=502, detail="Aucune entrée générée — réessaie ou reformule.")
|
||||
|
||||
name = str(parsed.get("name") or body.description).strip()[:120] or "Table générée"
|
||||
return GenerateTableResponseDTO(
|
||||
name=name,
|
||||
description=str(parsed.get("description") or "").strip(),
|
||||
entries=entries,
|
||||
)
|
||||
|
||||
|
||||
class ImproviseRollRequestDTO(BaseModel):
|
||||
table_name: str
|
||||
result_label: str
|
||||
result_detail: str = Field(default="")
|
||||
context: str = Field(default="")
|
||||
|
||||
|
||||
class ImproviseRollResponseDTO(BaseModel):
|
||||
narration: str
|
||||
|
||||
|
||||
@app.post("/improvise/table-roll", response_model=ImproviseRollResponseDTO)
|
||||
async def improvise_table_roll(
|
||||
body: ImproviseRollRequestDTO,
|
||||
llm: Annotated[LLMProvider, Depends(get_llm_provider)],
|
||||
) -> ImproviseRollResponseDTO:
|
||||
"""Brode un court récit (2-3 phrases) sur un résultat tiré, pour lancer la scène."""
|
||||
detail = f" ({body.result_detail.strip()})" if body.result_detail.strip() else ""
|
||||
context_block = f"\nContexte : {body.context.strip()}" if body.context.strip() else ""
|
||||
prompt = (
|
||||
"Tu es le Maître du Jeu. Les joueurs viennent de tirer sur la table "
|
||||
f"« {body.table_name.strip()} » et ont obtenu : « {body.result_label.strip()} »{detail}."
|
||||
f"{context_block}\n\n"
|
||||
"Décris en 2-3 phrases vivantes et immédiates ce qui se passe, pour lancer la scène. "
|
||||
"Pas de méta, pas d'options : juste la narration, en français."
|
||||
)
|
||||
try:
|
||||
raw = await llm.generate(prompt, temperature=0.8)
|
||||
except LLMProviderError as exc:
|
||||
raise HTTPException(status_code=502, detail=str(exc)) from exc
|
||||
return ImproviseRollResponseDTO(narration=raw.strip())
|
||||
|
||||
|
||||
# --- Notebooks (atelier RAG) : indexation des sources + chat ancré ----------
|
||||
|
||||
|
||||
class IndexSourceResponseDTO(BaseModel):
|
||||
chunks: int
|
||||
page_count: int
|
||||
ocr_page_count: int
|
||||
|
||||
|
||||
@app.post("/index/notebook-source", response_model=IndexSourceResponseDTO)
|
||||
async def index_notebook_source(
|
||||
rag: Annotated[NotebookRagUseCase, Depends(get_notebook_rag_use_case)],
|
||||
source_id: str = Form(...),
|
||||
file: UploadFile = File(...),
|
||||
) -> IndexSourceResponseDTO:
|
||||
"""Indexe une source PDF (extraction + embeddings + stockage vectoriel)."""
|
||||
content = await file.read()
|
||||
if not content:
|
||||
raise HTTPException(status_code=422, detail="Fichier PDF vide.")
|
||||
if len(content) > _MAX_PDF_BYTES:
|
||||
raise HTTPException(
|
||||
status_code=413, detail=f"PDF trop volumineux (> {_MAX_PDF_BYTES // (1024 * 1024)} Mo).")
|
||||
try:
|
||||
recap = await rag.index_source(source_id, content)
|
||||
except PdfExtractionError as exc:
|
||||
raise HTTPException(status_code=422, detail=str(exc)) from exc
|
||||
except EmbeddingError as exc:
|
||||
raise HTTPException(status_code=502, detail=str(exc)) from exc
|
||||
return IndexSourceResponseDTO(**recap)
|
||||
|
||||
|
||||
@app.delete("/index/notebook-source/{source_id}")
|
||||
def delete_notebook_source(source_id: str) -> dict[str, str]:
|
||||
"""Supprime les vecteurs d'une source (au DELETE d'une source/notebook)."""
|
||||
vector_store.delete(source_id)
|
||||
return {"status": "deleted", "source_id": source_id}
|
||||
|
||||
|
||||
class NotebookChatMessageDTO(BaseModel):
|
||||
role: str
|
||||
content: str
|
||||
|
||||
|
||||
class NotebookChatRequestDTO(BaseModel):
|
||||
source_ids: list[str] = Field(default_factory=list)
|
||||
messages: list[NotebookChatMessageDTO] = Field(default_factory=list)
|
||||
context: str = Field(default="")
|
||||
|
||||
|
||||
@app.post("/chat/notebook/stream")
|
||||
async def chat_notebook_stream(
|
||||
body: NotebookChatRequestDTO,
|
||||
use_case: Annotated[NotebookChatUseCase, Depends(get_notebook_chat_use_case)],
|
||||
settings: Annotated[Settings, Depends(get_settings)],
|
||||
) -> StreamingResponse:
|
||||
"""Chat ANCRÉ sur les sources (RAG) : récupère les passages pertinents puis
|
||||
streame la réponse. Évènements SSE : `token` {token}, `done` {}, `error` {message}."""
|
||||
messages = [ChatMessage(role=m.role, content=m.content) for m in body.messages]
|
||||
top_k = max(1, min(settings.rag_top_k, 200))
|
||||
|
||||
def _sse(event: str, data: dict) -> str:
|
||||
return f"event: {event}\ndata: {json.dumps(data, ensure_ascii=False)}\n\n"
|
||||
|
||||
async def event_stream() -> AsyncIterator[str]:
|
||||
try:
|
||||
async for token in use_case.stream(body.source_ids, messages, context=body.context, top_k=top_k):
|
||||
if token:
|
||||
yield _sse("token", {"token": token})
|
||||
yield _sse("done", {})
|
||||
except (LLMProviderError, EmbeddingError) as exc:
|
||||
yield _sse("error", {"message": str(exc)})
|
||||
except Exception as exc: # noqa: BLE001 — filet : pas de coupure brutale du flux.
|
||||
logger.exception("Chat notebook : erreur inattendue.")
|
||||
yield _sse("error", {"message": f"Erreur inattendue du Brain : {type(exc).__name__} : {exc}"})
|
||||
|
||||
return StreamingResponse(event_stream(), media_type="text/event-stream")
|
||||
|
||||
|
||||
@app.post("/chat/notebook/deep/stream")
|
||||
async def chat_notebook_deep_stream(
|
||||
body: NotebookChatRequestDTO,
|
||||
use_case: Annotated[NotebookDeepUseCase, Depends(get_notebook_deep_use_case)],
|
||||
) -> StreamingResponse:
|
||||
"""Analyse APPROFONDIE (map-reduce sur tout le document). Évènements SSE :
|
||||
`progress` {current,total} pendant la lecture, puis `token` {token}, puis `done`."""
|
||||
question = next((m.content for m in reversed(body.messages) if m.role == "user"), "")
|
||||
|
||||
def _sse(event: str, data: dict) -> str:
|
||||
return f"event: {event}\ndata: {json.dumps(data, ensure_ascii=False)}\n\n"
|
||||
|
||||
async def event_stream() -> AsyncIterator[str]:
|
||||
if not question.strip():
|
||||
yield _sse("error", {"message": "Question vide."})
|
||||
return
|
||||
try:
|
||||
async for ev in use_case.stream(body.source_ids, question, context=body.context):
|
||||
ev_type = ev.pop("type")
|
||||
yield _sse(ev_type, ev)
|
||||
except (LLMProviderError, EmbeddingError) as exc:
|
||||
yield _sse("error", {"message": str(exc)})
|
||||
except Exception as exc: # noqa: BLE001 — filet : pas de coupure brutale.
|
||||
logger.exception("Analyse approfondie : erreur inattendue.")
|
||||
yield _sse("error", {"message": f"Erreur inattendue du Brain : {type(exc).__name__} : {exc}"})
|
||||
|
||||
return StreamingResponse(event_stream(), media_type="text/event-stream")
|
||||
|
||||
|
||||
# --- Mapping DTO → domaine (frontière HTTP) ---------------------------------
|
||||
|
||||
|
||||
@@ -890,7 +1239,7 @@ class SettingsDTO(BaseModel):
|
||||
Les secrets (onemin_api_key) sont masques en lecture.
|
||||
"""
|
||||
|
||||
llm_provider: Literal["ollama", "onemin", "openrouter"]
|
||||
llm_provider: Literal["ollama", "onemin", "openrouter", "mistral", "gemini"]
|
||||
ollama_base_url: str
|
||||
llm_model: str
|
||||
onemin_model: str
|
||||
@@ -899,6 +1248,18 @@ class SettingsDTO(BaseModel):
|
||||
openrouter_model: str
|
||||
# True si une cle OpenRouter est deja configuree (cle elle-meme jamais renvoyee).
|
||||
openrouter_api_key_set: bool
|
||||
mistral_model: str
|
||||
# True si une cle Mistral est deja configuree (cle elle-meme jamais renvoyee).
|
||||
mistral_api_key_set: bool
|
||||
gemini_model: str
|
||||
# True si une cle Gemini est deja configuree (cle elle-meme jamais renvoyee).
|
||||
gemini_api_key_set: bool
|
||||
# Embeddings (RAG des ateliers) : provider + modeles + auto-pull Ollama.
|
||||
embedding_provider: Literal["ollama", "mistral"]
|
||||
ollama_embedding_model: str
|
||||
mistral_embedding_model: str
|
||||
auto_pull_embedding_model: bool
|
||||
rag_top_k: int
|
||||
# Fenetre de contexte effective passee au modele (num_ctx Ollama) — sert
|
||||
# aussi de plafond a la jauge de contexte UI.
|
||||
llm_num_ctx: int
|
||||
@@ -911,7 +1272,7 @@ class SettingsDTO(BaseModel):
|
||||
class SettingsUpdateDTO(BaseModel):
|
||||
"""Patch partiel des settings. Tous les champs sont optionnels."""
|
||||
|
||||
llm_provider: Literal["ollama", "onemin", "openrouter"] | None = None
|
||||
llm_provider: Literal["ollama", "onemin", "openrouter", "mistral", "gemini"] | None = None
|
||||
ollama_base_url: str | None = None
|
||||
llm_model: str | None = None
|
||||
onemin_model: str | None = None
|
||||
@@ -919,6 +1280,15 @@ class SettingsUpdateDTO(BaseModel):
|
||||
onemin_api_key: str | None = None
|
||||
openrouter_model: str | None = None
|
||||
openrouter_api_key: str | None = None
|
||||
mistral_model: str | None = None
|
||||
mistral_api_key: str | None = None
|
||||
gemini_model: str | None = None
|
||||
gemini_api_key: str | None = None
|
||||
embedding_provider: Literal["ollama", "mistral"] | None = None
|
||||
ollama_embedding_model: str | None = None
|
||||
mistral_embedding_model: str | None = None
|
||||
auto_pull_embedding_model: bool | None = None
|
||||
rag_top_k: int | None = None
|
||||
llm_num_ctx: int | None = None
|
||||
import_chunk_tokens: int | None = None
|
||||
llm_timeout_seconds: int | None = None
|
||||
@@ -933,6 +1303,15 @@ def _to_settings_dto(s: Settings) -> SettingsDTO:
|
||||
onemin_api_key_set=bool(s.onemin_api_key),
|
||||
openrouter_model=s.openrouter_model,
|
||||
openrouter_api_key_set=bool(s.openrouter_api_key),
|
||||
mistral_model=s.mistral_model,
|
||||
mistral_api_key_set=bool(s.mistral_api_key),
|
||||
gemini_model=s.gemini_model,
|
||||
gemini_api_key_set=bool(s.gemini_api_key),
|
||||
embedding_provider=s.embedding_provider,
|
||||
ollama_embedding_model=s.ollama_embedding_model,
|
||||
mistral_embedding_model=s.mistral_embedding_model,
|
||||
auto_pull_embedding_model=s.auto_pull_embedding_model,
|
||||
rag_top_k=s.rag_top_k,
|
||||
llm_num_ctx=s.llm_num_ctx,
|
||||
import_chunk_tokens=s.import_chunk_tokens,
|
||||
llm_timeout_seconds=s.llm_timeout_seconds,
|
||||
@@ -1138,6 +1517,99 @@ async def list_openrouter_models() -> dict[str, list[dict[str, object]]]:
|
||||
return {"models": models}
|
||||
|
||||
|
||||
# Repli statique si la cle Mistral n'est pas (encore) configuree ou si l'API est
|
||||
# injoignable — l'utilisateur peut quand meme choisir un modele. Liste curee
|
||||
# (juin 2026) ; pour l'extraction de PDF, prefere `large` (fidele, 128k) ou `small`.
|
||||
_MISTRAL_FALLBACK_MODELS = [
|
||||
"mistral-large-latest",
|
||||
"mistral-medium-latest",
|
||||
"mistral-small-latest",
|
||||
"open-mistral-nemo",
|
||||
"ministral-8b-latest",
|
||||
"ministral-3b-latest",
|
||||
"magistral-medium-latest",
|
||||
"magistral-small-latest",
|
||||
"pixtral-large-latest",
|
||||
"codestral-latest",
|
||||
]
|
||||
|
||||
|
||||
@app.get("/models/mistral")
|
||||
async def list_mistral_models(
|
||||
settings: Annotated[Settings, Depends(get_settings)],
|
||||
) -> dict[str, list[dict[str, object]]]:
|
||||
"""Catalogue des modeles Mistral. Dynamique si une cle est configuree
|
||||
(GET /v1/models, qui requiert l'auth), sinon repli statique.
|
||||
|
||||
Renvoie {models: [{id}]} (tous accessibles sur le tier gratuit Experiment)."""
|
||||
key = settings.mistral_api_key
|
||||
if not key:
|
||||
return {"models": [{"id": m} for m in _MISTRAL_FALLBACK_MODELS]}
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=20) as client:
|
||||
response = await client.get(
|
||||
"https://api.mistral.ai/v1/models",
|
||||
headers={"Authorization": f"Bearer {key}"},
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
except httpx.HTTPError:
|
||||
# Cle invalide / API down : on ne casse pas l'UI, on propose le repli.
|
||||
return {"models": [{"id": m} for m in _MISTRAL_FALLBACK_MODELS]}
|
||||
|
||||
ids = sorted({str(m.get("id")) for m in data.get("data", []) or [] if m.get("id")})
|
||||
if not ids:
|
||||
ids = _MISTRAL_FALLBACK_MODELS
|
||||
return {"models": [{"id": i} for i in ids]}
|
||||
|
||||
|
||||
# Repli statique Gemini (juin 2026). Pour l'extraction, prefere un Flash a grand
|
||||
# contexte ; `gemini-2.0-flash` a le quota gratuit le plus genereux.
|
||||
_GEMINI_FALLBACK_MODELS = [
|
||||
"gemini-2.0-flash",
|
||||
"gemini-2.0-flash-lite",
|
||||
"gemini-2.5-flash",
|
||||
"gemini-2.5-flash-lite",
|
||||
"gemini-2.5-pro",
|
||||
"gemini-1.5-flash",
|
||||
"gemini-1.5-pro",
|
||||
]
|
||||
|
||||
|
||||
@app.get("/models/gemini")
|
||||
async def list_gemini_models(
|
||||
settings: Annotated[Settings, Depends(get_settings)],
|
||||
) -> dict[str, list[dict[str, object]]]:
|
||||
"""Catalogue des modeles Gemini. Dynamique si une cle est configuree (endpoint
|
||||
OpenAI-compatible /openai/models), sinon repli statique. Renvoie {models:[{id}]}."""
|
||||
key = settings.gemini_api_key
|
||||
if not key:
|
||||
return {"models": [{"id": m} for m in _GEMINI_FALLBACK_MODELS]}
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=20) as client:
|
||||
response = await client.get(
|
||||
"https://generativelanguage.googleapis.com/v1beta/openai/models",
|
||||
headers={"Authorization": f"Bearer {key}"},
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
except httpx.HTTPError:
|
||||
return {"models": [{"id": m} for m in _GEMINI_FALLBACK_MODELS]}
|
||||
|
||||
# Les ids peuvent arriver prefixes "models/" → on nettoie pour que la valeur
|
||||
# selectionnee soit directement utilisable dans l'appel chat. On garde les
|
||||
# modeles "gemini-*" (hors embeddings/aqa) pour ne pas noyer la liste.
|
||||
ids: set[str] = set()
|
||||
for m in data.get("data", []) or []:
|
||||
mid = str(m.get("id") or "")
|
||||
if mid.startswith("models/"):
|
||||
mid = mid[len("models/"):]
|
||||
if mid.startswith("gemini-"):
|
||||
ids.add(mid)
|
||||
clean = sorted(ids) if ids else _GEMINI_FALLBACK_MODELS
|
||||
return {"models": [{"id": i} for i in clean]}
|
||||
|
||||
|
||||
@app.get("/models/onemin")
|
||||
def list_onemin_models() -> dict[str, list[dict[str, object]]]:
|
||||
"""Catalogue statique des modeles 1min.ai, groupes par fournisseur.
|
||||
|
||||
Reference in New Issue
Block a user