File size: 5,413 Bytes
8200ee1
 
 
 
 
 
 
 
 
 
b666236
 
8200ee1
02f31d3
8200ee1
b301ad2
b666236
 
 
 
 
8200ee1
02f31d3
8200ee1
b666236
1d72461
b666236
 
 
 
8200ee1
 
b301ad2
 
 
 
 
 
 
 
 
 
 
8200ee1
 
 
 
 
 
 
 
02f31d3
8200ee1
 
 
 
 
 
 
 
 
 
 
b301ad2
8200ee1
 
 
02f31d3
 
 
 
 
8200ee1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
02f31d3
8200ee1
 
 
 
 
 
 
b301ad2
02f31d3
8200ee1
 
 
 
 
 
 
 
 
b301ad2
02f31d3
8200ee1
 
 
 
b301ad2
8200ee1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b666236
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8200ee1
 
 
b301ad2
02f31d3
8200ee1
 
 
 
 
 
 
 
 
 
b301ad2
02f31d3
8200ee1
02f31d3
8200ee1
 
 
 
 
 
 
 
 
02f31d3
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
"""
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)