DavidJyes's picture
Update app.py
16a091c verified
Raw
History Blame
7.82 kB
from fastapi import FastAPI #, HTTPException
from pydantic import BaseModel, Field
import joblib
import json
import numpy as np
from datetime import datetime
from typing import Optional
import os
# Initialisation
app = FastAPI(
title="API Détection de Fraude",
description="API de détection de fraude dans les transactions bancaires",
version="1.0.0"
)
# Chargement du modèle et des encodeurs
model = joblib.load('fraud_model.pkl')
le_category = joblib.load('le_category.pkl')
le_gender = joblib.load('le_gender.pkl')
le_state = joblib.load('le_state.pkl')
with open('mappings.json', 'r') as f:
mappings = json.load(f)
# Modèle de données
class Transaction(BaseModel):
amt: float = Field(..., description="Montant de la transaction", example=125.50)
category: str = Field(..., description="Catégorie du marchand", example="personal_care")
merchant: str = Field(..., description="Nom du marchand", example="fraud_Kirlin and Sons")
trans_date_trans_time: str = Field(..., description="Date et heure", example="2020-06-21 12:14:25")
gender: str = Field(..., description="Genre (M/F)", example="M")
state: str = Field(..., description="État (code à 2 lettres)", example="SC")
lat: float = Field(..., description="Latitude du client", example=33.9659)
long: float = Field(..., description="Longitude du client", example=-80.9355)
city_pop: int = Field(..., description="Population de la ville", example=333497)
dob: str = Field(..., description="Date de naissance", example="1968-03-19")
merch_lat: float = Field(..., description="Latitude du marchand", example=33.986391)
merch_long: float = Field(..., description="Longitude du marchand", example=-81.200714)
cc_num: Optional[str] = Field(None, description="Numéro de carte (optionnel)")
avg_amt: Optional[float] = Field(50.0, description="Montant moyen historique")
std_amt: Optional[float] = Field(30.0, description="Écart-type historique")
nb_trans: Optional[int] = Field(10, description="Nombre de transactions historiques")
class PredictionResponse(BaseModel):
is_fraud: bool
fraud_probability: float
risk_level: str
details: dict
# Routes
@app.get("/")
def read_root():
return {
"message": "API de Détection de Fraude",
"version": "1.0.0",
"endpoints": {
"/predict": "POST - Prédire une transaction",
"/health": "GET - Statut de l'API",
"/categories": "GET - Liste des catégories",
"/docs": "GET - Documentation interactive"
}
}
@app.get("/health")
def health_check():
return {
"status": "healthy"},
# "model_loaded": model is not None,
#"timestamp": datetime.now().isoformat()
# }
@app.get("/categories")
def get_categories():
return {
"categories": list(mappings['categories'].keys()),
"states": list(mappings['states'].keys()),
"genders": list(mappings['genders'].keys())
}
@app.post("/predict", response_model=PredictionResponse)
def predict_fraud(transaction: Transaction):
try:
# Parsing des dates
trans_dt = datetime.strptime(transaction.trans_date_trans_time, "%Y-%m-%d %H:%M:%S")
dob_dt = datetime.strptime(transaction.dob, "%Y-%m-%d")
# Features temporelles
hour = trans_dt.hour
day_of_week = trans_dt.weekday()
day = trans_dt.day
month = trans_dt.month
age = (trans_dt - dob_dt).days // 365
# Distance
distance = np.sqrt(
(transaction.lat - transaction.merch_lat)**2 +
(transaction.long - transaction.merch_long)**2
) * 111
# Encodage
if transaction.category not in mappings['categories']:
raise HTTPException(400, f"Catégorie inconnue: {transaction.category}")
if transaction.gender not in mappings['genders']:
raise HTTPException(400, f"Genre inconnu: {transaction.gender}")
if transaction.state not in mappings['states']:
raise HTTPException(400, f"État inconnu: {transaction.state}")
category_encoded = mappings['categories'][transaction.category]
gender_encoded = mappings['genders'][transaction.gender]
state_encoded = mappings['states'][transaction.state]
# Construction du vecteur de features
features = np.array([[
transaction.amt,
hour,
day_of_week,
day,
month,
age,
category_encoded,
gender_encoded,
state_encoded,
transaction.lat,
transaction.long,
transaction.city_pop,
distance,
transaction.avg_amt,
transaction.std_amt,
transaction.nb_trans
]])
# Prédiction
fraud_proba = model.predict_proba(features)[0][1]
is_fraud = fraud_proba > 0.5
# Niveau de risque
if fraud_proba < 0.3:
risk_level = "Faible"
elif fraud_proba < 0.7:
risk_level = "Moyen"
else:
risk_level = "Élevé"
return PredictionResponse(
is_fraud=bool(is_fraud),
fraud_probability=float(fraud_proba),
risk_level=risk_level,
details={
"montant": transaction.amt,
"categorie": transaction.category,
"heure": hour,
"age_client": age,
"distance_km": round(distance, 2)
}
)
except ValueError as e:
raise HTTPException(status_code=400, detail=f"Erreur de format: {str(e)}")
except Exception as e:
raise HTTPException(status_code=500, detail=f"Erreur interne: {str(e)}")
#if __name__ == "__main__":
# import uvicorn
# uvicorn.run(app, host="0.0.0.0", port=7860)
# if __name__ == "__main__":
# import gradio as gr
# from app import predict_fraud, Transaction # si tu es déjà dans app.py, inutile de réimporter
# # Wrapper pour adapter les inputs de Gradio au modèle FastAPI
# def api_predict_wrapper(
# amt, category, merchant, trans_date_trans_time, gender, state,
# lat, long, city_pop, dob, merch_lat, merch_long,
# cc_num=None, avg_amt=50, std_amt=30, nb_trans=10
# ):
# tx = Transaction(
# amt=amt,
# category=category,
# merchant=merchant,
# trans_date_trans_time=trans_date_trans_time,
# gender=gender,
# state=state,
# lat=lat,
# long=long,
# city_pop=city_pop,
# dob=dob,
# merch_lat=merch_lat,
# merch_long=merch_long,
# cc_num=cc_num,
# avg_amt=avg_amt,
# std_amt=std_amt,
# nb_trans=nb_trans
# )
# result = predict_fraud(tx)
# return result.dict()
# # Définition de l'interface Gradio
# iface = gr.Interface(
# fn=api_predict_wrapper,
# inputs=[
# gr.Number(label="Montant"),
# gr.Textbox(label="Catégorie"),
# gr.Textbox(label="Marchand"),
# gr.Textbox(label="Date/Heure"),
# gr.Textbox(label="Genre"),
# gr.Textbox(label="État"),
# gr.Number(label="Latitude"),
# gr.Number(label="Longitude"),
# gr.Number(label="Population ville"),
# gr.Textbox(label="Date de naissance"),
# gr.Number(label="Lat march."),
# gr.Number(label="Long march."),
# ],
# outputs=gr.JSON(label="Résultat")
# )
# # Lancement du front Gradio
# iface.launch()