ThePredictors / app.py
YBDixit's picture
Update app.py
44d9192 verified
Raw
History Blame Contribute Delete
32.5 kB
"""
RAG Chat API - Gustave Eiffel Hackathon 2026
=============================================
Base : app_candis.py
Ajout : question_type externe (champ QueryRequest) pour honorer le type
fourni par l'évaluateur HF sans changer l'architecture Candis.
"""
import os
import re
import json
import logging
import re
import time
from pathlib import Path
from typing import Optional
os.environ.setdefault("ANONYMIZED_TELEMETRY", "False")
import requests as http_requests
import gradio as gr
from fastapi import FastAPI, HTTPException
from fastapi.responses import JSONResponse
from pydantic import BaseModel
import chromadb
from chromadb.config import Settings
from langchain_text_splitters import RecursiveCharacterTextSplitter
from pypdf import PdfReader
from llm import call_llm as call_llm_with_metrics
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
logging.getLogger("chromadb.telemetry.product.posthog").setLevel(logging.CRITICAL)
# ---------------------------------------------------------------------------
# Configuration
# ---------------------------------------------------------------------------
DATA_DIR = Path("/data") if Path("/data").is_dir() else Path("./data")
CHROMA_PERSIST_DIR = str(DATA_DIR / "chroma_db")
TRAIN_DOCS_DIR = Path("./train_data")
COLLECTION_NAME = "rag_documents"
CHUNK_SIZE = 900
CHUNK_OVERLAP = 100
TOP_K_RESULTS = 3
DENSE_CANDIDATES = 20
KEYWORD_CANDIDATES = 20
MAX_FINAL_CONTEXTS = 6
MINI_MODEL_CONFIDENCE_THRESHOLD = 0.35
EMBEDDING_BATCH_SIZE = 32
_CONFIG_PATH = DATA_DIR / "config.json"
if not _CONFIG_PATH.exists():
_CONFIG_PATH = Path(__file__).parent / "config.json"
logger.warning(
f"No config.json found in {DATA_DIR} — falling back to root config.json."
)
with open(_CONFIG_PATH, encoding="utf-8") as _f:
_config = json.load(_f)
EMBEDDING_ENDPOINT_URL = _config["embedding"]["endpoint_url"]
EMBEDDING_MODEL_NAME = _config["embedding"]["model"]
LLM_ENDPOINT_URL = _config["llm"]["endpoint_url"]
LLM_MODEL_NAME = _config["llm"]["model"]
LLM_MINI_MODEL_NAME = _config["llm"].get("mini_model", "gpt-5-mini")
LLM_MINI_ENDPOINT_URL = _config["llm"].get(
"mini_endpoint_url",
LLM_ENDPOINT_URL.replace(LLM_MODEL_NAME, LLM_MINI_MODEL_NAME),
)
LLM_MAX_TOKENS = _config["llm"].get("max_completion_tokens", 512)
LLM_TEMPERATURE = _config["llm"].get("temperature", 0.7)
LLM_TOP_P = _config["llm"].get("top_p", 0.95)
AZURE_API_KEY = os.environ.get("AZURE_API_KEY")
if not AZURE_API_KEY:
logger.warning("AZURE_API_KEY is not set — LLM and embedding calls will fail.")
_PROMPT_TEMPLATE_PATH = Path(__file__).parent / "prompts" / "rag_prompt.txt"
RAG_PROMPT_TEMPLATE = _PROMPT_TEMPLATE_PATH.read_text(encoding="utf-8")
logger.info(f"Embedding model configured: {EMBEDDING_MODEL_NAME} via Azure OpenAI")
logger.info(f"Initializing ChromaDB at: {CHROMA_PERSIST_DIR}")
chroma_client = chromadb.PersistentClient(
path=CHROMA_PERSIST_DIR,
settings=Settings(anonymized_telemetry=False),
)
collection = chroma_client.get_or_create_collection(
name=COLLECTION_NAME,
metadata={"hnsw:space": "cosine"},
)
logger.info(f"ChromaDB collection '{COLLECTION_NAME}' ready. Documents: {collection.count()}")
logger.info(f"LLM configured: {LLM_MODEL_NAME} via {LLM_ENDPOINT_URL}")
_KEYWORD_CACHE: dict = {"count": -1, "items": []}
STOPWORDS = {
"alors", "avec", "avoir", "dans", "dont", "elle", "elles", "entre", "etre",
"leur", "leurs", "mais", "nous", "pour", "plus", "quel", "quelle", "quelles",
"quels", "sans", "sont", "tout", "tous", "une", "vous", "the", "and", "for",
"that", "this", "with", "from", "what", "which", "aux", "des", "les", "est",
"sur", "par", "ses", "ces", "qui", "que", "quoi", "comment", "pourquoi",
}
QUESTION_PROFILES = {
"definition": {
"top_k": 3,
"max_context_chars": 900,
"guidance": "Definition question: answer in 1-3 factual sentences. Do not add a broad analysis.",
},
"formula": {
"top_k": 4,
"max_context_chars": 1100,
"guidance": "Formula/calculation question: give the exact formula or calculation rule when present, then one short clarification.",
},
"comparison": {
"top_k": 5,
"max_context_chars": 1100,
"guidance": "Comparison question: compare the concepts in short bullet points or a compact paragraph, covering all requested sides.",
},
"explanation": {
"top_k": 5,
"max_context_chars": 1100,
"guidance": "Explanation question: provide a concise but complete explanation, usually 3-5 sentences.",
},
"regulatory": {
"top_k": 5,
"max_context_chars": 1200,
"guidance": "Regulatory question: be precise, avoid overgeneralizing, and mention the source/page evidence.",
},
"source_specific": {
"top_k": 4,
"max_context_chars": 1200,
"guidance": "Source-specific question: focus on the document or domain named in the question.",
},
"qcm": {
"top_k": 2,
"max_context_chars": 350,
"max_total_context_chars": 900,
"retrieval_mode": "dense",
"guidance": "Multiple-choice question: compare the options, select the best-supported letter, and keep the justification short.",
},
"synthesis": {
"top_k": 6,
"max_context_chars": 1000,
"guidance": "Synthesis/analysis question: use diverse sources and structure the answer in short, focused points.",
},
"default": {
"top_k": TOP_K_RESULTS,
"max_context_chars": 1000,
"guidance": "General question: answer as briefly as possible while preserving correctness and useful detail.",
},
}
# ---------------------------------------------------------------------------
# Helper Functions
# ---------------------------------------------------------------------------
def is_multiple_choice_question(query: str) -> bool:
has_a = re.search(r"\bA\s*[\)\.\:\-]", query, re.IGNORECASE)
has_b = re.search(r"\bB\s*[\)\.\:\-]", query, re.IGNORECASE)
has_c = re.search(r"\bC\s*[\)\.\:\-]", query, re.IGNORECASE)
has_d = re.search(r"\bD\s*[\)\.\:\-]", query, re.IGNORECASE)
q = query.lower()
return bool(has_a and has_b and has_c and has_d) or "qcm" in q
def clean_text(text: str) -> str:
text = text.replace("\x00", " ")
text = re.sub(r"[ \t]+", " ", text)
text = re.sub(r"\n{3,}", "\n\n", text)
return text.strip()
def extract_page_number(text: str) -> Optional[int]:
match = re.search(r"\[Page\s+(\d+)\]", text)
return int(match.group(1)) if match else None
def tokenize_keywords(text: str) -> list[str]:
tokens = re.findall(r"\b[\wÀ-ÖØ-öø-ÿ.-]{2,}\b", text.lower())
keywords = []
for token in tokens:
normalized = token.strip("._-")
if len(normalized) < 2 or normalized in STOPWORDS:
continue
keywords.append(normalized)
return keywords
def is_qcm_query(query: str) -> bool:
lowered = query.lower()
qcm_markers = [
"qcm", "choisir", "choisissez", "quelle proposition", "option",
"a.", "b.", "c.", "d.", "réponse a", "réponse b", "réponse c", "réponse d",
"reponse a", "reponse b", "reponse c", "reponse d",
]
return any(marker in lowered for marker in qcm_markers) or bool(
re.search(r"(^|\n)\s*[abcd][\).:-]\s+", lowered)
)
def is_complex_query(query: str) -> bool:
lowered = query.lower()
complex_markers = [
"compare", "explique", "pourquoi", "limites", "avantages",
"inconvénients", "inconvenients", "dans quelle mesure", "synthèse",
"synthese", "analyse",
]
return len(query.split()) > 28 or any(marker in lowered for marker in complex_markers)
def classify_question(query: str) -> str:
lowered = query.lower().strip()
if is_qcm_query(query):
return "qcm"
if re.search(r"\b(compare|diff[eé]rence|distingue|oppos[ée]|versus| vs |scr et mcr)\b", lowered):
return "comparison"
if re.search(r"\b(formule|calcul|calcule|calculer|montant|ratio|taux|probabilit[eé])\b", lowered):
return "formula"
if re.search(r"\b(r[eé]glement|directive|solvabilit[eé]\s*ii|ifrs\s*17|norme|exigence)\b", lowered):
return "regulatory"
if re.search(r"\b(m[eé]moire|document|pdf|source|dans le|dans la|selon)\b", lowered):
return "source_specific"
if re.search(r"\b(synth[eè]se|synth[eé]tise|analyse|r[eé]sume|principaux|globalement)\b", lowered):
return "synthesis"
if re.search(r"\b(pourquoi|comment|explique|impact|effet|limites|avantages|inconv[eé]nients)\b", lowered):
return "explanation"
if re.search(r"^(c[' ]?est quoi|qu[' ]?est[- ]ce que|d[eé]finis|d[eé]finition|que signifie)\b", lowered):
return "definition"
return "default"
def get_question_profile(query: str, question_type_override: Optional[str] = None) -> dict:
"""Résout le profil. Si question_type_override est fourni (ex: par l'évaluateur HF), il prime."""
if question_type_override is not None:
qt = question_type_override.upper().replace("-", "_").replace(" ", "_")
if qt in ("QCM", "MCQ", "MULTIPLE_CHOICE"):
question_type = "qcm"
else:
question_type = classify_question(query)
else:
question_type = classify_question(query)
profile = dict(QUESTION_PROFILES.get(question_type, QUESTION_PROFILES["default"]))
profile["type"] = question_type
return profile
def choose_final_top_k(query: str, requested_top_k: Optional[int] = None, question_type_override: Optional[str] = None) -> int:
profile = get_question_profile(query, question_type_override)
if requested_top_k is not None:
return max(1, min(requested_top_k, MAX_FINAL_CONTEXTS))
return max(1, min(profile["top_k"], MAX_FINAL_CONTEXTS))
def get_keyword_items() -> list[dict]:
count = collection.count()
if count == 0:
return []
if _KEYWORD_CACHE["count"] == count:
return _KEYWORD_CACHE["items"]
results = collection.get(include=["documents", "metadatas"])
items = []
for idx, text in enumerate(results.get("documents") or []):
metadata = (results.get("metadatas") or [{}])[idx] or {}
source = metadata.get("source", "unknown")
chunk_index = metadata.get("chunk_index", idx)
page = metadata.get("page")
items.append({
"text": text,
"source": source,
"chunk_index": chunk_index,
"page": page,
"keywords": set(tokenize_keywords(text)),
})
_KEYWORD_CACHE["count"] = count
_KEYWORD_CACHE["items"] = items
return items
def score_keyword_matches(query: str, item: dict) -> float:
query_terms = tokenize_keywords(query)
if not query_terms:
return 0.0
item_terms = item.get("keywords") or set(tokenize_keywords(item.get("text", "")))
score = 0.0
for term in query_terms:
if term in item_terms:
score += 1.0
elif len(term) > 4 and any(term in candidate for candidate in item_terms):
score += 0.45
return score / max(len(set(query_terms)), 1)
def keyword_retrieve(query: str, limit: int = KEYWORD_CANDIDATES) -> list[dict]:
scored = []
for item in get_keyword_items():
keyword_score = score_keyword_matches(query, item)
if keyword_score <= 0:
continue
scored.append({
"text": item["text"],
"source": item["source"],
"page": item.get("page"),
"chunk_index": item["chunk_index"],
"similarity_score": 0.0,
"keyword_score": keyword_score,
"retrieval_method": "keyword",
})
scored.sort(key=lambda ctx: ctx["keyword_score"], reverse=True)
return scored[:limit]
def extract_text_from_pdf(pdf_path: Path) -> str:
reader = PdfReader(str(pdf_path))
pages_text = []
for page_num, page in enumerate(reader.pages, start=1):
text = page.extract_text()
if text and text.strip():
cleaned = clean_text(text)
if cleaned:
pages_text.append(f"[Page {page_num}]\n{cleaned}")
full_text = "\n\n".join(pages_text)
logger.info(
f"Extracted {len(reader.pages)} pages from PDF: "
f"{pdf_path.name} ({len(full_text)} chars)"
)
return full_text
def chunk_text(text: str, source: str = "unknown") -> list[dict]:
cleaned_text = clean_text(text)
splitter = RecursiveCharacterTextSplitter(
chunk_size=CHUNK_SIZE,
chunk_overlap=CHUNK_OVERLAP,
separators=["\n\n", "\n", ". ", " ", ""],
)
chunks = splitter.split_text(cleaned_text)
documents = []
for i, chunk in enumerate(chunks):
chunk = clean_text(chunk)
if len(chunk) < 80:
continue
documents.append({
"text": chunk,
"source": source,
"page": extract_page_number(chunk),
"chunk_index": i,
})
return documents
def _generate_embeddings_batch(texts: list[str]) -> list[list[float]]:
headers = {"api-key": AZURE_API_KEY, "Content-Type": "application/json"}
payload = {"input": texts, "model": EMBEDDING_MODEL_NAME}
resp = http_requests.post(EMBEDDING_ENDPOINT_URL, headers=headers, json=payload, timeout=120)
resp.raise_for_status()
data = resp.json()
return [item["embedding"] for item in data["data"]]
def generate_embeddings(texts: list[str]) -> list[list[float]]:
all_embeddings = []
try:
for i in range(0, len(texts), EMBEDDING_BATCH_SIZE):
batch = texts[i:i + EMBEDDING_BATCH_SIZE]
all_embeddings.extend(_generate_embeddings_batch(batch))
if len(texts) > EMBEDDING_BATCH_SIZE:
time.sleep(0.2)
return all_embeddings
except http_requests.exceptions.HTTPError as e:
logger.error(f"Embedding API call failed: {e}")
raise HTTPException(status_code=503, detail=f"Embedding service unavailable: {str(e)}")
except (http_requests.exceptions.JSONDecodeError, ValueError) as e:
logger.error(f"Embedding API returned invalid JSON: {e}")
raise HTTPException(status_code=502, detail="Embedding service returned an invalid response")
except (KeyError, IndexError) as e:
logger.error(f"Unexpected embedding response format: {e}")
raise HTTPException(status_code=502, detail="Unexpected response from embedding service")
def add_documents_to_vectorstore(documents: list[dict]) -> int:
if not documents:
return 0
texts = [doc["text"] for doc in documents]
embeddings = generate_embeddings(texts)
existing_count = collection.count()
ids = [f"doc_{existing_count + i}" for i in range(len(documents))]
metadatas = [
{
"source": doc["source"],
"page": doc.get("page") or -1,
"chunk_index": doc["chunk_index"],
}
for doc in documents
]
collection.add(ids=ids, embeddings=embeddings, documents=texts, metadatas=metadatas)
logger.info(f"Added {len(documents)} chunks to vector store. Total: {collection.count()}")
_KEYWORD_CACHE["count"] = -1
return len(documents)
def merge_and_rerank_contexts(query: str, contexts: list[dict], final_k: int) -> list[dict]:
merged: dict[tuple[str, int], dict] = {}
for ctx in contexts:
key = (ctx.get("source", "unknown"), ctx.get("chunk_index", -1))
existing = merged.get(key)
if existing is None:
merged[key] = ctx
continue
existing["similarity_score"] = max(
existing.get("similarity_score", 0.0), ctx.get("similarity_score", 0.0),
)
existing["keyword_score"] = max(
existing.get("keyword_score", 0.0), ctx.get("keyword_score", 0.0),
)
methods = set(str(existing.get("retrieval_method", "")).split("+"))
methods.update(str(ctx.get("retrieval_method", "")).split("+"))
existing["retrieval_method"] = "+".join(sorted(m for m in methods if m))
reranked = []
source_counts: dict[str, int] = {}
for ctx in merged.values():
similarity = max(ctx.get("similarity_score", 0.0), 0.0)
keyword = max(ctx.get("keyword_score", 0.0), score_keyword_matches(query, ctx))
text_len = len(ctx.get("text", ""))
length_bonus = 0.05 if 250 <= text_len <= CHUNK_SIZE + 300 else 0.0
page_bonus = 0.03 if ctx.get("page", -1) not in (None, -1) else 0.0
source = ctx.get("source", "unknown")
diversity_penalty = min(source_counts.get(source, 0) * 0.08, 0.24)
source_counts[source] = source_counts.get(source, 0) + 1
combined_score = (
0.65 * similarity
+ 0.35 * keyword
+ length_bonus
+ page_bonus
- diversity_penalty
)
ctx["keyword_score"] = keyword
ctx["combined_score"] = combined_score
reranked.append(ctx)
reranked.sort(key=lambda ctx: ctx["combined_score"], reverse=True)
return reranked[:final_k]
def retrieve_relevant_context(
query: str,
top_k: int = TOP_K_RESULTS,
question_type_override: Optional[str] = None,
) -> list[dict]:
if collection.count() == 0:
return []
profile = get_question_profile(query, question_type_override)
final_k = choose_final_top_k(query, top_k, question_type_override)
query_embedding = generate_embeddings([query])[0]
dense_results = final_k if profile.get("retrieval_mode") == "dense" else max(DENSE_CANDIDATES, final_k)
results = collection.query(
query_embeddings=[query_embedding],
n_results=min(dense_results, collection.count()),
include=["documents", "metadatas", "distances"],
)
contexts = []
for i in range(len(results["documents"][0])):
metadata = results["metadatas"][0][i] or {}
contexts.append({
"text": results["documents"][0][i],
"source": metadata.get("source", "unknown"),
"page": metadata.get("page", -1),
"chunk_index": metadata.get("chunk_index", i),
"similarity_score": 1 - results["distances"][0][i],
"keyword_score": 0.0,
"retrieval_method": "dense",
})
if profile.get("retrieval_mode") == "dense":
return contexts[:final_k]
contexts.extend(keyword_retrieve(query, limit=KEYWORD_CANDIDATES))
return merge_and_rerank_contexts(query, contexts, final_k=final_k)
def truncate_context_text(text: str, max_chars: int) -> str:
text = clean_text(text)
if len(text) <= max_chars:
return text
return text[:max_chars].rsplit(" ", 1)[0].strip() + "..."
def build_rag_prompt(query: str, contexts: list[dict], profile: dict) -> str:
max_context_chars = profile.get("max_context_chars", 1000)
if profile.get("type") == "qcm":
blocks = []
total_chars = 0
max_total_chars = profile.get("max_total_context_chars", 900)
for idx, ctx in enumerate(contexts, start=1):
block = f"[C{idx}]\n{truncate_context_text(ctx['text'], max_context_chars)}"
if total_chars + len(block) > max_total_chars:
break
blocks.append(block)
total_chars += len(block)
context_text = "\n\n".join(blocks)
else:
context_text = "\n\n".join(
f"[Source: {ctx['source']} | Page: {ctx.get('page', -1)}]\n"
f"{truncate_context_text(ctx['text'], max_context_chars)}"
for ctx in contexts
)
prompt = RAG_PROMPT_TEMPLATE.format(
context=context_text,
question=query,
question_type=profile.get("type", "default"),
response_guidance=profile.get("guidance", QUESTION_PROFILES["default"]["guidance"]),
)
return prompt
def choose_llm_for_query(query: str, contexts: list[dict]) -> tuple[str, str, str]:
if not contexts:
return LLM_MINI_MODEL_NAME, LLM_MINI_ENDPOINT_URL, "no_context"
best_score = max(ctx.get("combined_score", 0.0) for ctx in contexts)
avg_score = sum(ctx.get("combined_score", 0.0) for ctx in contexts) / len(contexts)
if is_qcm_query(query):
return LLM_MINI_MODEL_NAME, LLM_MINI_ENDPOINT_URL, "qcm_mini"
if best_score < MINI_MODEL_CONFIDENCE_THRESHOLD or avg_score < MINI_MODEL_CONFIDENCE_THRESHOLD / 2:
return LLM_MODEL_NAME, LLM_ENDPOINT_URL, "low_retrieval_confidence"
if is_complex_query(query):
return LLM_MODEL_NAME, LLM_ENDPOINT_URL, "complex_query"
return LLM_MINI_MODEL_NAME, LLM_MINI_ENDPOINT_URL, "default_mini"
def is_insufficient_context_answer(answer: str) -> bool:
lowered = (answer or "").lower()
markers = [
"contexte insuffisant",
"insufficient context",
"i don't have enough information",
"je n'ai pas assez",
"je ne dispose pas",
]
return any(marker in lowered for marker in markers)
def rag_query(query: str, top_k: Optional[int] = None, question_type: Optional[str] = None) -> dict:
start_time = time.perf_counter()
profile = get_question_profile(query, question_type)
contexts = retrieve_relevant_context(query, top_k=top_k, question_type_override=question_type)
if not contexts:
elapsed_ms = round((time.perf_counter() - start_time) * 1000, 2)
return {
"answer": "Les documents ne permettent pas de répondre.",
"sources": [],
"explanation": "Aucun document n'a été trouvé dans la base vectorielle.",
"total_token": 0,
"prompt_tokens": 0,
"completion_tokens": 0,
"cached_tokens": 0,
"co2_grams": None,
"energy_kwh": None,
"run_time_in_ms": elapsed_ms,
}
is_qcm = is_multiple_choice_question(query)
best_score = max(ctx["similarity_score"] for ctx in contexts)
# Anti-hallucination uniquement pour les questions ouvertes.
# Pour les QCM, on laisse le prompt gérer :
# - option "information non disponible" si elle existe ;
# - sinon réponse libre "Les documents ne permettent pas de répondre."
if not is_qcm and best_score < 0.45:
elapsed_ms = round((time.perf_counter() - start_time) * 1000, 2)
return {
"answer": "Les documents ne permettent pas de répondre.",
"sources": [
{
"source": ctx["source"],
"score": ctx["similarity_score"],
"ref_text": ctx["text"],
}
for ctx in contexts
],
"explanation": "Le contexte récupéré est trop peu pertinent pour répondre de manière fiable.",
"total_token": 0,
"prompt_tokens": 0,
"completion_tokens": 0,
"cached_tokens": 0,
"co2_grams": None,
"energy_kwh": None,
"run_time_in_ms": elapsed_ms,
}
prompt = build_rag_prompt(query, contexts, profile)
selected_model, selected_endpoint, routing_reason = choose_llm_for_query(query, contexts)
llm_result = call_llm_with_metrics(
prompt,
endpoint_url=selected_endpoint,
api_key=AZURE_API_KEY,
model=selected_model,
max_completion_tokens=LLM_MAX_TOKENS,
temperature=LLM_TEMPERATURE,
top_p=LLM_TOP_P,
)
raw_content = llm_result["content"]
tokens = llm_result["tokens"]
total_token = tokens["total"]
json_str = raw_content.strip()
if json_str.startswith("```"):
json_str = json_str.split("\n", 1)[-1]
json_str = json_str.rsplit("```", 1)[0].strip()
try:
parsed = json.loads(json_str)
answer = parsed["answer"]
explanation = parsed["explanation"]
except (json.JSONDecodeError, KeyError):
answer = raw_content
explanation = "LLM did not return a structured explanation."
insufficient_context = is_insufficient_context_answer(answer)
if insufficient_context:
answer = "Contexte insuffisant"
explanation = ""
elif profile["type"] == "qcm":
explanation = ""
elapsed_ms = round((time.perf_counter() - start_time) * 1000, 2)
return {
"answer": answer,
"sources": [
{
"source": ctx["source"],
"page": ctx.get("page", -1),
"score": ctx.get("combined_score", ctx.get("similarity_score", 0.0)),
"dense_score": ctx.get("similarity_score", 0.0),
"keyword_score": ctx.get("keyword_score", 0.0),
"retrieval_method": ctx.get("retrieval_method", "unknown"),
"ref_text": ctx["text"],
}
for ctx in contexts
],
"explanation": explanation,
"total_token": total_token,
"prompt_tokens": tokens["prompt"],
"completion_tokens": tokens["completion"],
"cached_tokens": tokens["cached"],
"co2_grams": llm_result["co2_grams"],
"energy_kwh": llm_result["energy_kwh"],
"run_time_in_ms": elapsed_ms,
"selected_model": selected_model,
"routing_reason": routing_reason,
"question_type": profile["type"],
"insufficient_context": insufficient_context,
}
# ---------------------------------------------------------------------------
# Ingest Train Documents
# ---------------------------------------------------------------------------
def ingest_train_documents():
if collection.count() > 0:
logger.info("Vector store already has documents, skipping ingestion.")
return
if not TRAIN_DOCS_DIR.exists():
logger.warning(f"No train_data directory found at: {TRAIN_DOCS_DIR}")
return
for file_path in TRAIN_DOCS_DIR.rglob("*.txt"):
logger.info(f"Ingesting text file: {file_path.name}")
text = file_path.read_text(encoding="utf-8")
chunks = chunk_text(text, source=file_path.name)
add_documents_to_vectorstore(chunks)
for file_path in TRAIN_DOCS_DIR.rglob("*.pdf"):
logger.info(f"Ingesting PDF file: {file_path.name}")
text = extract_text_from_pdf(file_path)
if text.strip():
chunks = chunk_text(text, source=file_path.name)
add_documents_to_vectorstore(chunks)
else:
logger.warning(f"No extractable text found in: {file_path.name}")
logger.info(f"Train document ingestion complete. Total chunks: {collection.count()}")
# ---------------------------------------------------------------------------
# FastAPI Application
# ---------------------------------------------------------------------------
app = FastAPI(
title="RAG Chat API - Gustave Eiffel Hackathon 2026",
description="A RAG system with /query endpoint for evaluation",
version="1.0.0",
)
class QueryRequest(BaseModel):
query: str
top_k: Optional[int] = None
question_type: Optional[str] = None
class IngestRequest(BaseModel):
text: str
source: str = "user_upload"
@app.post("/query")
async def query_endpoint(request: QueryRequest):
logger.info(f"Query received: {request.query!r} (question_type={request.question_type})")
result = rag_query(request.query, top_k=request.top_k, question_type=request.question_type)
return JSONResponse(content=result)
@app.post("/ingest")
async def ingest_endpoint(request: IngestRequest):
chunks = chunk_text(request.text, source=request.source)
count = add_documents_to_vectorstore(chunks)
return JSONResponse(
content={
"status": "success",
"chunks_added": count,
"total_chunks": collection.count(),
}
)
@app.get("/health")
async def health_check():
return {
"status": "healthy",
"documents_in_store": collection.count(),
"embedding_model": EMBEDDING_MODEL_NAME,
"llm_model": LLM_MODEL_NAME,
"mini_llm_model": LLM_MINI_MODEL_NAME,
"chunk_size": CHUNK_SIZE,
"chunk_overlap": CHUNK_OVERLAP,
"top_k_results": TOP_K_RESULTS,
}
# ---------------------------------------------------------------------------
# Gradio UI
# ---------------------------------------------------------------------------
def gradio_query(question: str) -> tuple[str, str, str, str, str]:
if not question.strip():
return "Please enter a question.", "", "", "", ""
result = rag_query(question)
if result.get("insufficient_context"):
answer = result["answer"]
else:
sources_text = "\n".join(
f" - {s['source']} (relevance: {s['score']:.2f})" for s in result["sources"][:3]
)
answer = f"{result['answer']}\n\n📚 Sources:\n{sources_text}" if result["sources"] else result["answer"]
explanation = result.get("explanation", "")
token_info = str(result.get("total_token", 0))
co2_value = result.get("co2_grams")
co2_info = f"{co2_value:.4f} g" if isinstance(co2_value, (int, float)) else "N/A"
run_time = f"{result.get('run_time_in_ms', 0)} ms"
return answer, explanation, token_info, co2_info, run_time
def gradio_ingest(text: str, source_name: str) -> str:
if not text.strip():
return "Please provide text to ingest."
chunks = chunk_text(
text,
source=source_name or "user_upload",
)
count = add_documents_to_vectorstore(chunks)
return f"✅ Ingested {count} chunks. Total documents in store: {collection.count()}"
with gr.Blocks(title="RAG Chat API - Gustave Eiffel Hackathon") as demo:
gr.Markdown(
"""
# 🗼 RAG Chat API - Gustave Eiffel Hackathon 2026
This application demonstrates a complete **Retrieval-Augmented Generation (RAG)** system.
**API Endpoint:** Use `POST /query` with `{"query": "your question"}` for programmatic access.
---
"""
)
with gr.Tab("💬 Chat"):
gr.Markdown("Ask questions about the ingested documents.")
with gr.Row():
query_input = gr.Textbox(
label="Your Question",
placeholder="e.g., What is the Eiffel Tower made of?",
lines=2,
)
query_button = gr.Button("Ask", variant="primary")
query_output = gr.Textbox(label="Answer", lines=8, interactive=False)
query_explanation = gr.Textbox(label="Explanation", lines=3, interactive=False)
with gr.Row():
query_tokens = gr.Textbox(label="Total Tokens", interactive=False)
query_co2 = gr.Textbox(label="CO2 Emission", interactive=False)
query_runtime = gr.Textbox(label="Run Time", interactive=False)
query_button.click(
fn=gradio_query,
inputs=query_input,
outputs=[
query_output,
query_explanation,
query_tokens,
query_co2,
query_runtime,
],
)
with gr.Tab("📄 Ingest Documents"):
gr.Markdown("Add new documents to the knowledge base.")
doc_text = gr.Textbox(
label="Document Text",
placeholder="Paste your document text here...",
lines=10,
)
doc_source = gr.Textbox(
label="Source Name",
placeholder="e.g., my_document.txt",
value="user_upload",
)
ingest_button = gr.Button("Ingest Document", variant="primary")
ingest_output = gr.Textbox(label="Status", interactive=False)
ingest_button.click(
fn=gradio_ingest,
inputs=[doc_text, doc_source],
outputs=ingest_output,
)
with gr.Tab("ℹ️ API Info"):
gr.Markdown(
"""
## API Endpoints
### POST /query
```json
{
"query": "What is the Eiffel Tower?",
"top_k": 3,
"question_type": "QCM"
}
```
### GET /health
Returns system health and document count.
"""
)
app = gr.mount_gradio_app(app, demo, path="/")
if __name__ == "__main__":
import uvicorn
uvicorn.run(
app,
host="0.0.0.0",
port=7860,
)