| """ |
| 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 |
|
|
| |
|
|
| logging.basicConfig(level=logging.INFO) |
| logger = logging.getLogger(__name__) |
|
|
|
|
| |
|
|
| 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) |
|
|
|
|
| |
|
|
|
|
| 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), |
| } |
|
|
|
|
| |
|
|
|
|
| @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) |
|
|
| |
| 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())} |
|
|
|
|
| |
|
|
| if __name__ == "__main__": |
| import uvicorn |
|
|
| uvicorn.run("src.api.main:app", host="0.0.0.0", port=7860, reload=True) |
|
|