""" API de prédiction du churn employé. Contenu: - Configuration de l'application FastAPI - Fonction de prédiction - Définition des endpoints - Point d'entrée de l'application """ import logging from fastapi import FastAPI, HTTPException from fastapi.responses import Response from src.api.schemas import DonneesEmploye, ResultatPrediction from src.db.logger_db import ( log_api_operation, log_prediction_input, log_prediction_output, ) from src.utils.model_loader import get_modele from src.utils.transformer import transformer_donnees # CONFIGURATION DU LOGGER logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) # APPLICATION FASTAPI tags_metadata = [ { "name": "Service API", "description": "Endpoint du service API pour l'inférence.", }, { "name": "Infrastructure", "description": "Endpoints courants d'infrastructure pour l'observabilité (contrôle d'intégrité, métriques).", }, ] app = FastAPI( title="API Prédiction Churn Employé", description=""" API qui prédit si un employé va quitter l'entreprise. ## Fonctionnalités - Prédiction du churn (départ) d'un employé - Validation automatique des données entrantes - Documentation interactive Swagger UI (/docs) ## Données requises L'API attend les données complètes d'un employé incluant : - Informations personnelles (âge, genre, salaire, etc.) - Expérience professionnelle - Scores de satisfaction - Informations sur le poste """, version="1.0.0", docs_url="/docs", redoc_url="/redoc", openapi_tags=tags_metadata, ) @app.get("/favicon.ico", include_in_schema=False) async def favicon(): return Response(status_code=204) # FONCTIONS DE PRÉDICTION def predire_churn(donnees: dict) -> dict: """ Fonction principale de prédiction. Étapes: 1. Transformer les données (encodage, etc.) 2. Appliquer le modèle 3. Retourner le résultat Args: donnees: Dict avec les données de l'employé Returns: Dict avec prédiction, probabilité et classe Raises: HTTPException: Si les features ne correspondent pas au modèle """ donnees_transformees = transformer_donnees(donnees) modele = get_modele() try: prediction = modele.predict(donnees_transformees)[0] probabilites = modele.predict_proba(donnees_transformees)[0] except Exception as e: raise HTTPException( status_code=422, detail=( "Erreur de features: Les colonnes fournies ne correspondent pas " f"aux features du modèle entraîné. Message original: {e}" ), ) from e proba_depart = probabilites[1] return { "prediction": "Oui" if prediction == 1 else "Non", "probabilite": round(float(proba_depart), 3), "classe": int(prediction), } # ENDPOINTS (routes de l'API) @app.get("/", tags=["Infrastructure"]) async def racine(): """Endpoint racine - informations générales.""" return { "message": "API Prédiction Churn Employé", "version": "1.0.0", "documentation": "/docs", "redoc": "/redoc", } @app.get("/health", tags=["Infrastructure"]) async def health_check(): """Vérifie que l'API fonctionne.""" return {"status": "ok", "message": "L'API est opérationnelle"} @app.post("/predire", response_model=ResultatPrediction, tags=["Service API"]) def predire(employe: DonneesEmploye): """ Endpoint de prédiction du churn. Args: employe: Données de l'employé (validées par Pydantic) Returns: Prédiction avec probabilité de départ Example: ```json { "prediction": "Oui", "probabilite": 0.75, "classe": 1 } ``` """ donnees = employe.model_dump(exclude_none=True) resultat = predire_churn(donnees) # Logging dans la base de données (optionnel, ne bloque pas l'API) input_id = log_prediction_input(donnees) log_prediction_output( input_id=input_id, prediction=resultat["prediction"], probabilite=resultat["probabilite"], classe=resultat["classe"], ) log_api_operation( operation="PREDICT", table_cible="prediction_inputs", details=f"id_employee={donnees.get('id_employee')}, prediction={resultat['prediction']}", statut="SUCCESS", ) return ResultatPrediction(**resultat) @app.get("/modele/info", tags=["Infrastructure"]) async def infos_modele(): """Retourne des informations sur le modèle.""" return { "type": "RandomForestClassifier", "description": "Modèle de prédiction du churn employé", "version": "1.0.0", "nombre_features": 24, "cible": "Prédiction du départ d'un employé (0=restera, 1=partira)", } @app.get("/modele/features", tags=["Infrastructure"]) async def liste_features(): """Retourne la liste des features attendues par le modèle.""" from src.utils.transformer import get_liste_features return {"features": get_liste_features(), "nombre": len(get_liste_features())} # POINT D'ENTRÉE if __name__ == "__main__": import uvicorn uvicorn.run("src.api.main:app", host="0.0.0.0", port=7860, reload=True)