JeremyMikaleff's picture
feat: add document retrieval and prompt enrichment pipeline
84da485
Raw
History Blame Contribute Delete
12.9 kB
# ------------------------------------------------------ Imports -------------------------------------------------------------------------------|
# Standard library imports
# Third-party library imports
from fastapi import FastAPI, HTTPException, status
from fastapi.openapi.docs import get_swagger_ui_html
from fastapi.responses import HTMLResponse
from fastapi.staticfiles import StaticFiles
from pydantic import BaseModel, Field
from sqlalchemy import inspect
# Local application imports
from database import check_database_connection, engine
from routers.collections import router as collections_router
from routers.documents import router as documents_router
from routers.prompts import router as prompts_router
# ----------------------------------------------------------------------------------------------------------------------------------------------|
# ------------------------------------------------------ Schémas des réponses Swagger ----------------------------------------------------------|
class ApplicationResponse(BaseModel):
"""Informations générales sur l'application."""
application: str = Field(examples=["Anderson — Semantic Retrieval for RAG"], )
version: str = Field(examples=["0.2.0"], )
status: str = Field(examples=["running"], )
documentation: str = Field(examples=["/docs"], )
class HealthResponse(BaseModel):
"""État général de l'API."""
status: str = Field(examples=["ok"], )
class DatabaseHealthResponse(BaseModel):
"""État de la connexion PostgreSQL."""
status: str = Field(examples=["ok"], )
database: str = Field(examples=["connected"], )
connection_test: bool = Field(examples=[True], )
class DatabaseColumnResponse(BaseModel):
"""Description d'une colonne PostgreSQL."""
name: str
type: str
nullable: bool
primary_key: bool
foreign_key: str | None = None
class DatabaseTableResponse(BaseModel):
"""Description d'une table PostgreSQL."""
name: str
columns: list[DatabaseColumnResponse]
class DatabaseSchemaResponse(BaseModel):
"""Description du schéma public de la base de données."""
database: str = Field(examples=["PostgreSQL"], )
schema_name: str = Field(examples=["public"], )
table_count: int = Field(examples=[4], )
tables: list[DatabaseTableResponse]
# ----------------------------------------------------------------------------------------------------------------------------------------------|
# ------------------------------------------------------ Métadonnées Swagger -------------------------------------------------------------------|
swagger_tags = [
{
"name": "API Health",
"description": (
"Endpoints permettant de vérifier la disponibilité de l'API "
"et de ses dépendances."
),
},
{
"name": "Database",
"description": (
"Consultation technique de la base PostgreSQL utilisée "
"par la démonstration Anderson."
),
},
{
"name": "Collections",
"description": "Création, consultation, modification et suppression des collections documentaires.",
},
{
"name": "Documents",
"description": "Upload, extraction, vectorisation, consultation et suppression des documents.",
},
{
"name": "Prompts",
"description": "Recherche vectorielle et génération de prompts enrichis à partir des documents.",
},
]
# ----------------------------------------------------------------------------------------------------------------------------------------------|
# ------------------------------------------------------ Création de l'application FastAPI -----------------------------------------------------|
app = FastAPI(
title="Anderson — Semantic Retrieval for RAG",
description=(
"![Logo Anderson](/static/anderson-logo.png)\n\n"
"## Démonstration Hugging Face Spaces\n\n"
"Cette version publique et allégée d'Anderson permet d'organiser des documents en collections, "
"d'extraire et découper leur contenu, de calculer des embeddings, puis d'effectuer une recherche "
"vectorielle afin de produire un prompt enrichi.\n\n"
"Le code de cette démonstration est consultable directement dans le "
"[dépôt du Space Hugging Face]"
"(https://huggingface.co/spaces/JeremyMikaleff/Anderson-Semantic_Retrieval_for_RAG/tree/main).\n\n"
"## Projet Anderson original\n\n"
"Le [projet complet disponible sur GitHub]"
"(https://github.com/Projets-finaux-Simplon-2024/API-Anderson-BDD) "
"présente une architecture plus étendue comprenant :\n\n"
"- des tests automatisés de l'API et de ses services ;\n"
"- une authentification avec JWT Bearer Token ;\n"
"- une gestion des utilisateurs, des rôles et des permissions ;\n"
"- PostgreSQL avec l'extension pgvector ;\n"
"- MinIO pour assurer la persistance des documents uploadés ;\n"
"- MLflow pour le registre, le versionnement et le monitoring du modèle d'embedding ;\n"
"- le suivi de métriques de similarité cosinus ;\n"
"- une étude comparative de trois modèles d'embedding : "
"`OrdalieTech/Solon-embeddings-large-0.1`, `bge-m3-custom-fr` et "
"`Cohere embed-multilingual-v3.0`.\n\n"
"L'architecture de la base avait été préparée pour stocker les embeddings des trois modèles. "
"L'intégration opérationnelle finale s'est concentrée sur Solon, afin de conserver un pipeline "
"exécutable avec les ressources matérielles disponibles."
),
version="0.2.0",
docs_url=None,
redoc_url="/redoc",
openapi_url="/openapi.json",
openapi_tags=swagger_tags,
swagger_ui_parameters={
"defaultModelsExpandDepth": -1,
"docExpansion": "none",
"displayRequestDuration": True,
"filter": True,
},
)
app.include_router(collections_router)
app.include_router(documents_router)
app.include_router(prompts_router)
app.mount("/static", StaticFiles(directory="static"), name="static")
# ----------------------------------------------------------------------------------------------------------------------------------------------|
# ------------------------------------------------------ Documentation Swagger personnalisée ---------------------------------------------------|
@app.get("/docs", include_in_schema=False)
def custom_swagger_ui():
swagger_page = get_swagger_ui_html(
openapi_url=app.openapi_url,
title=f"{app.title} — Documentation",
swagger_favicon_url="static/favicon.png",
swagger_ui_parameters={
"defaultModelsExpandDepth": -1,
"docExpansion": "none",
"displayRequestDuration": True,
"filter": True,
},
)
html_content = swagger_page.body.decode("utf-8").replace(
"</head>",
'<link rel="stylesheet" href="/static/swagger-dark.css"></head>',
)
return HTMLResponse(html_content)
# ----------------------------------------------------------------------------------------------------------------------------------------------|
# ------------------------------------------------------ Endpoints API Health ------------------------------------------------------------------|
@app.get(
"/",
response_model=ApplicationResponse,
status_code=status.HTTP_200_OK,
summary="Informations sur l'application",
description=(
"Retourne le nom, la version et l'état général de l'application Anderson."
),
tags=["API Health"],
)
def root() -> ApplicationResponse:
return ApplicationResponse(
application="Anderson — Semantic Retrieval for RAG",
version=app.version,
status="running",
documentation="/docs",
)
@app.get(
"/health",
response_model=HealthResponse,
status_code=status.HTTP_200_OK,
summary="Vérifier la disponibilité de l'API",
description=(
"Vérifie que le serveur FastAPI est démarré et capable de répondre "
"aux requêtes HTTP."
),
tags=["API Health"],
)
def health() -> HealthResponse:
return HealthResponse(status="ok")
@app.get(
"/health/database",
response_model=DatabaseHealthResponse,
status_code=status.HTTP_200_OK,
summary="Vérifier la connexion PostgreSQL",
description=(
"Exécute une requête minimale sur la base Neon PostgreSQL afin de vérifier "
"que la connexion configurée avec DATABASE_URL est opérationnelle."
),
responses={
status.HTTP_503_SERVICE_UNAVAILABLE: {
"description": "La base PostgreSQL est indisponible.",
},
},
tags=["API Health"],
)
def database_health() -> DatabaseHealthResponse:
try:
connection_test = check_database_connection()
except Exception as exception:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="Database unavailable",
) from exception
return DatabaseHealthResponse(
status="ok",
database="connected",
connection_test=connection_test,
)
# ----------------------------------------------------------------------------------------------------------------------------------------------|
# ------------------------------------------------------ Endpoint schéma PostgreSQL ------------------------------------------------------------|
@app.get(
"/database/schema",
response_model=DatabaseSchemaResponse,
status_code=status.HTTP_200_OK,
summary="Consulter le schéma de la base de données",
description=(
"Retourne les tables, colonnes, types, clés primaires et clés étrangères "
"du schéma public PostgreSQL. Aucune donnée métier ni information de "
"connexion sensible n'est exposée."
),
responses={
status.HTTP_503_SERVICE_UNAVAILABLE: {
"description": "Le schéma PostgreSQL ne peut pas être consulté.",
},
},
tags=["Database"],
)
def get_database_schema() -> DatabaseSchemaResponse:
try:
database_inspector = inspect(engine)
table_names = sorted(
database_inspector.get_table_names(schema="public"),
)
tables: list[DatabaseTableResponse] = []
for table_name in table_names:
primary_key = database_inspector.get_pk_constraint(
table_name,
schema="public",
)
primary_key_columns = set(
primary_key.get("constrained_columns") or [],
)
foreign_keys = database_inspector.get_foreign_keys(
table_name,
schema="public",
)
foreign_key_mapping: dict[str, str] = {}
for foreign_key in foreign_keys:
constrained_columns = foreign_key.get(
"constrained_columns",
[],
)
referred_table = foreign_key.get("referred_table")
referred_columns = foreign_key.get(
"referred_columns",
[],
)
for column_name, referred_column in zip(
constrained_columns,
referred_columns,
):
foreign_key_mapping[column_name] = (
f"{referred_table}.{referred_column}"
)
columns: list[DatabaseColumnResponse] = []
for column in database_inspector.get_columns(
table_name,
schema="public",
):
column_name = column["name"]
columns.append(
DatabaseColumnResponse(
name=column_name,
type=str(column["type"]),
nullable=column["nullable"],
primary_key=column_name in primary_key_columns,
foreign_key=foreign_key_mapping.get(column_name),
)
)
tables.append(
DatabaseTableResponse(
name=table_name,
columns=columns,
)
)
except Exception as exception:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="Database schema unavailable",
) from exception
return DatabaseSchemaResponse(
database="PostgreSQL",
schema_name="public",
table_count=len(tables),
tables=tables,
)
# ----------------------------------------------------------------------------------------------------------------------------------------------|