Spaces:
Runtime error
Runtime error
Déploiement auto depuis GitHub Actions avec LFS
Browse files- .gitattributes +1 -0
- Dockerfile +36 -0
- app/__init__.py +0 -0
- app/main.py +128 -0
- app/model_loader.py +38 -0
- app/schemas.py +153 -0
- models/pmvl_catboost_final.cbm +3 -0
- models/pmvl_feature_columns.txt +22 -0
- requirements.txt +8 -0
- tests/__init__.py +0 -0
- tests/test_api.py +85 -0
.gitattributes
CHANGED
|
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
*.cbm filter=lfs diff=lfs merge=lfs -text
|
Dockerfile
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
FROM python:3.12-slim
|
| 2 |
+
|
| 3 |
+
# 1. Création d'un utilisateur standard (UID 1000 requis par Hugging Face)
|
| 4 |
+
RUN useradd -m -u 1000 user
|
| 5 |
+
USER user
|
| 6 |
+
|
| 7 |
+
# 2. Variables d'environnement pour l'utilisateur et Python
|
| 8 |
+
ENV HOME=/home/user \
|
| 9 |
+
PATH=/home/user/.local/bin:$PATH \
|
| 10 |
+
PYTHONDONTWRITEBYTECODE=1 \
|
| 11 |
+
PYTHONUNBUFFERED=1
|
| 12 |
+
|
| 13 |
+
# 3. Répertoire de travail
|
| 14 |
+
WORKDIR $HOME/app
|
| 15 |
+
|
| 16 |
+
# 4. Copie des dépendances avec les bons droits
|
| 17 |
+
COPY --chown=user requirements.txt .
|
| 18 |
+
|
| 19 |
+
# 5. Mise à jour de pip (pour enlever le 'notice') et installation
|
| 20 |
+
RUN pip install --no-cache-dir --upgrade pip && \
|
| 21 |
+
pip install --no-cache-dir -r requirements.txt
|
| 22 |
+
|
| 23 |
+
# 6. Copie du code et des modèles
|
| 24 |
+
COPY --chown=user app/ ./app/
|
| 25 |
+
COPY --chown=user models/ ./models/
|
| 26 |
+
|
| 27 |
+
# 7. Variables de l'API avec le nouveau chemin
|
| 28 |
+
ENV GLOBAL_THRESHOLD=0.45
|
| 29 |
+
ENV MODEL_PATH=$HOME/app/models/pmvl_catboost_final.cbm
|
| 30 |
+
ENV MODEL_FEATURES_PATH=$HOME/app/models/pmvl_feature_columns.txt
|
| 31 |
+
|
| 32 |
+
# 8. Exposer le port 7860 (Port par défaut de Hugging Face Spaces)
|
| 33 |
+
EXPOSE 7860
|
| 34 |
+
|
| 35 |
+
# 9. Démarrage de l'API sur le port 7860
|
| 36 |
+
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "7860"]
|
app/__init__.py
ADDED
|
File without changes
|
app/main.py
ADDED
|
@@ -0,0 +1,128 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from fastapi import FastAPI, HTTPException
|
| 2 |
+
from fastapi.middleware.cors import CORSMiddleware
|
| 3 |
+
import os
|
| 4 |
+
import pandas as pd
|
| 5 |
+
import numpy as np
|
| 6 |
+
|
| 7 |
+
from .schemas import PMVLFeatures, PredictionResponse
|
| 8 |
+
from .model_loader import get_model, get_feature_columns
|
| 9 |
+
|
| 10 |
+
GLOBAL_THRESHOLD_ENV = "GLOBAL_THRESHOLD"
|
| 11 |
+
DEFAULT_THRESHOLD = 0.45
|
| 12 |
+
|
| 13 |
+
app = FastAPI(
|
| 14 |
+
title="API de prédiction de la qualité PMVL",
|
| 15 |
+
description="API exposant le modèle CatBoost pour estimer la précision des PMVL.",
|
| 16 |
+
version="1.0.0",
|
| 17 |
+
)
|
| 18 |
+
|
| 19 |
+
app.add_middleware(
|
| 20 |
+
CORSMiddleware,
|
| 21 |
+
allow_origins=["*"],
|
| 22 |
+
allow_credentials=True,
|
| 23 |
+
allow_methods=["*"],
|
| 24 |
+
allow_headers=["*"],
|
| 25 |
+
)
|
| 26 |
+
|
| 27 |
+
GROUP_KEYS = [
|
| 28 |
+
"PMVL[ENTITE]",
|
| 29 |
+
"PMVL[Selected Fund code]",
|
| 30 |
+
"PMVL[ISIN]",
|
| 31 |
+
"PMVL[Ref Unik Asset]",
|
| 32 |
+
]
|
| 33 |
+
|
| 34 |
+
def make_position_group(df: pd.DataFrame) -> pd.Series:
|
| 35 |
+
"""Recrée la colonne position_group comme dans le notebook"""
|
| 36 |
+
existing = [c for c in GROUP_KEYS if c in df.columns]
|
| 37 |
+
if not existing:
|
| 38 |
+
return pd.Series(
|
| 39 |
+
np.arange(len(df)).astype(str), index=df.index, name="position_group"
|
| 40 |
+
)
|
| 41 |
+
return (
|
| 42 |
+
df[existing]
|
| 43 |
+
.astype(str)
|
| 44 |
+
.fillna("NA")
|
| 45 |
+
.agg("||".join, axis=1)
|
| 46 |
+
.rename("position_group")
|
| 47 |
+
)
|
| 48 |
+
|
| 49 |
+
def prepare_catboost_features(X: pd.DataFrame, cat_cols: list):
|
| 50 |
+
"""Prépare les features (remplace None par 'MISSING' pour les cat, etc.)"""
|
| 51 |
+
X = X.copy()
|
| 52 |
+
for c in X.columns:
|
| 53 |
+
if c in cat_cols:
|
| 54 |
+
# Remplacement des valeurs manquantes par 'MISSING' et conversion en string
|
| 55 |
+
X[c] = X[c].fillna("MISSING").astype(str)
|
| 56 |
+
else:
|
| 57 |
+
# Remplacement par NaN et forçage numérique
|
| 58 |
+
X[c] = pd.to_numeric(X[c], errors="coerce")
|
| 59 |
+
|
| 60 |
+
for c in X.columns:
|
| 61 |
+
if X[c].dtype == bool:
|
| 62 |
+
X[c] = X[c].astype(int)
|
| 63 |
+
return X
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
@app.on_event("startup")
|
| 67 |
+
def load_model_on_startup():
|
| 68 |
+
try:
|
| 69 |
+
get_model()
|
| 70 |
+
get_feature_columns()
|
| 71 |
+
except Exception as e:
|
| 72 |
+
print(f"Erreur critique lors du chargement initial : {e}")
|
| 73 |
+
|
| 74 |
+
@app.get("/health", tags=["diagnostic"])
|
| 75 |
+
def health_check():
|
| 76 |
+
return {"status": "ok", "message": "L'API PMVL est opérationnelle."}
|
| 77 |
+
|
| 78 |
+
@app.post("/predict", response_model=PredictionResponse, tags=["prédiction"])
|
| 79 |
+
def predict_pmvl(features: PMVLFeatures):
|
| 80 |
+
try:
|
| 81 |
+
model = get_model()
|
| 82 |
+
feature_columns = get_feature_columns()
|
| 83 |
+
|
| 84 |
+
# 1) Récupérer les données brutes
|
| 85 |
+
raw_dict = features.model_dump(by_alias=True)
|
| 86 |
+
raw_dict.pop("PMVL[Holding date]", None)
|
| 87 |
+
|
| 88 |
+
# 2) Construire un DataFrame avec TOUTES les colonnes attendues
|
| 89 |
+
row = {col: np.nan for col in feature_columns}
|
| 90 |
+
for k, v in raw_dict.items():
|
| 91 |
+
if k in row and v is not None:
|
| 92 |
+
row[k] = v
|
| 93 |
+
|
| 94 |
+
raw_df = pd.DataFrame([row])
|
| 95 |
+
|
| 96 |
+
# 3) Recréer position_group si elle fait partie des features
|
| 97 |
+
if "position_group" in feature_columns:
|
| 98 |
+
raw_df["position_group"] = make_position_group(raw_df)
|
| 99 |
+
|
| 100 |
+
# 4) S'assurer de l'ordre exact des colonnes
|
| 101 |
+
df_input = raw_df[feature_columns]
|
| 102 |
+
|
| 103 |
+
# 5) Appliquer le nettoyage en demandant AU MODÈLE quelles sont les colonnes catégorielles !
|
| 104 |
+
cat_indices = model.get_cat_feature_indices()
|
| 105 |
+
cat_cols = [feature_columns[i] for i in cat_indices]
|
| 106 |
+
|
| 107 |
+
df_input = prepare_catboost_features(df_input, cat_cols)
|
| 108 |
+
|
| 109 |
+
# 6) Prédire
|
| 110 |
+
proba = float(model.predict_proba(df_input)[:, 1][0])
|
| 111 |
+
|
| 112 |
+
threshold = float(os.getenv(GLOBAL_THRESHOLD_ENV, DEFAULT_THRESHOLD))
|
| 113 |
+
prediction_bool = bool(proba >= threshold)
|
| 114 |
+
|
| 115 |
+
return PredictionResponse(
|
| 116 |
+
proba_bonne_estimation=proba,
|
| 117 |
+
prediction=prediction_bool,
|
| 118 |
+
seuil_applique=threshold,
|
| 119 |
+
fund_code=features.fund_code,
|
| 120 |
+
ref_unik_asset=features.ref_unik_asset,
|
| 121 |
+
)
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
except Exception as e:
|
| 125 |
+
raise HTTPException(
|
| 126 |
+
status_code=500,
|
| 127 |
+
detail=f"Erreur lors du traitement de la prédiction : {e}",
|
| 128 |
+
)
|
app/model_loader.py
ADDED
|
@@ -0,0 +1,38 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
from functools import lru_cache
|
| 3 |
+
from catboost import CatBoostClassifier
|
| 4 |
+
|
| 5 |
+
# Variables d'environnement pour rendre le chemin dynamique (utile pour Docker)
|
| 6 |
+
MODEL_PATH_ENV = "MODEL_PATH"
|
| 7 |
+
DEFAULT_MODEL_PATH = "models/pmvl_catboost_final.cbm"
|
| 8 |
+
|
| 9 |
+
FEATURES_PATH_ENV = "FEATURES_PATH"
|
| 10 |
+
DEFAULT_FEATURES_PATH = "models/pmvl_feature_columns.txt"
|
| 11 |
+
|
| 12 |
+
@lru_cache(maxsize=1)
|
| 13 |
+
def get_model() -> CatBoostClassifier:
|
| 14 |
+
"""
|
| 15 |
+
Charge le modèle CatBoost depuis le disque.
|
| 16 |
+
Grâce à @lru_cache, cette fonction ne s'exécute réellement qu'une seule fois.
|
| 17 |
+
"""
|
| 18 |
+
model_path = os.getenv(MODEL_PATH_ENV, DEFAULT_MODEL_PATH)
|
| 19 |
+
|
| 20 |
+
if not os.path.exists(model_path):
|
| 21 |
+
raise FileNotFoundError(f"Le fichier du modèle est introuvable au chemin : {model_path}")
|
| 22 |
+
|
| 23 |
+
print("Chargement du modèle CatBoost en mémoire...")
|
| 24 |
+
model = CatBoostClassifier()
|
| 25 |
+
model.load_model(model_path)
|
| 26 |
+
print("Modèle chargé avec succès !")
|
| 27 |
+
|
| 28 |
+
return model
|
| 29 |
+
|
| 30 |
+
#chargeur pour la liste de colonnes utilisées par le modèle (pour s'assurer que l'ordre des features est correct)
|
| 31 |
+
@lru_cache(maxsize=1)
|
| 32 |
+
def get_feature_columns() -> list[str]:
|
| 33 |
+
features_path = os.getenv(FEATURES_PATH_ENV, DEFAULT_FEATURES_PATH)
|
| 34 |
+
if not os.path.exists(features_path):
|
| 35 |
+
raise FileNotFoundError(f"Le fichier des features est introuvable au chemin : {features_path}")
|
| 36 |
+
with open(features_path, encoding="utf-8") as f:
|
| 37 |
+
cols = [line.strip() for line in f if line.strip()]
|
| 38 |
+
return cols
|
app/schemas.py
ADDED
|
@@ -0,0 +1,153 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from pydantic import BaseModel, Field, validator
|
| 2 |
+
from typing import Optional
|
| 3 |
+
from datetime import date
|
| 4 |
+
|
| 5 |
+
class PMVLFeatures(BaseModel):
|
| 6 |
+
"""
|
| 7 |
+
Schéma Pydantic représentant une observation d'entrée pour la prédiction PMVL.
|
| 8 |
+
Les noms des attributs sont des alias Python-friendly pour les colonnes complexes du DataFrame.
|
| 9 |
+
"""
|
| 10 |
+
# ------------------ Colonnes de type Date ------------------
|
| 11 |
+
holding_date: date = Field(
|
| 12 |
+
...,
|
| 13 |
+
alias="PMVL[Holding date]",
|
| 14 |
+
description="Date de la position (Trading/Holding date)"
|
| 15 |
+
)
|
| 16 |
+
quote_date: Optional[str] = Field(
|
| 17 |
+
None,
|
| 18 |
+
alias="PMVL[Quote Date]",
|
| 19 |
+
description="Date de la cotation (peut contenir des valeurs manquantes)"
|
| 20 |
+
)
|
| 21 |
+
|
| 22 |
+
# ------------------ Colonnes Numériques (float64) ------------------
|
| 23 |
+
pmvl_estim: Optional[float] = Field(
|
| 24 |
+
None,
|
| 25 |
+
alias="PMVL[PMVL Estimé]",
|
| 26 |
+
description="Estimation de la PMVL au temps t"
|
| 27 |
+
)
|
| 28 |
+
pfm_indice_perf: Optional[float] = Field(
|
| 29 |
+
None,
|
| 30 |
+
alias="PMVL[PFMIndice.Perf J / J-1]",
|
| 31 |
+
description="Performance de l'indice J / J-1"
|
| 32 |
+
)
|
| 33 |
+
prmp_pmvl: Optional[float] = Field(
|
| 34 |
+
None,
|
| 35 |
+
alias="PMVL[PRMP PMVL]",
|
| 36 |
+
description="Realized PMVL at t"
|
| 37 |
+
)
|
| 38 |
+
prmp_vnc: Optional[float] = Field(
|
| 39 |
+
None,
|
| 40 |
+
alias="PMVL[PRMP VNC]",
|
| 41 |
+
description="Realized VNC at t"
|
| 42 |
+
)
|
| 43 |
+
prmp_mtm: Optional[float] = Field(
|
| 44 |
+
None,
|
| 45 |
+
alias="PMVL[PRMP MtM]",
|
| 46 |
+
description="Realized MtM at t"
|
| 47 |
+
)
|
| 48 |
+
quantity: float = Field(
|
| 49 |
+
...,
|
| 50 |
+
alias="PMVL[Quantity]",
|
| 51 |
+
description="Quantité de l'actif"
|
| 52 |
+
)
|
| 53 |
+
purch_val_clean: float = Field(
|
| 54 |
+
...,
|
| 55 |
+
alias="PMVL[Purch. Val. (clean) (ptf cur.)]",
|
| 56 |
+
description="Valeur d'achat clean en devise du portefeuille"
|
| 57 |
+
)
|
| 58 |
+
quote: float = Field(
|
| 59 |
+
...,
|
| 60 |
+
alias="PMVL[Quote]",
|
| 61 |
+
description="Cotation de l'actif"
|
| 62 |
+
)
|
| 63 |
+
vnc_agrege_dirty: float = Field(
|
| 64 |
+
...,
|
| 65 |
+
alias="PMVL[VNC Agrege dirty (ptf cur.)]",
|
| 66 |
+
description="VNC Agrégée dirty en devise du portefeuille"
|
| 67 |
+
)
|
| 68 |
+
|
| 69 |
+
# ------------------ Colonnes Catégorielles (object) ------------------
|
| 70 |
+
entite: str = Field(
|
| 71 |
+
...,
|
| 72 |
+
alias="PMVL[ENTITE]",
|
| 73 |
+
description="Entité légale"
|
| 74 |
+
)
|
| 75 |
+
isin: str = Field(
|
| 76 |
+
...,
|
| 77 |
+
alias="PMVL[ISIN]",
|
| 78 |
+
description="Code ISIN de l'instrument"
|
| 79 |
+
)
|
| 80 |
+
orig_name: str = Field(
|
| 81 |
+
...,
|
| 82 |
+
alias="PMVL[Orig. name]",
|
| 83 |
+
description="Nom original de l'actif"
|
| 84 |
+
)
|
| 85 |
+
ticker: str = Field(
|
| 86 |
+
...,
|
| 87 |
+
alias="PMVL[Parametres_Indices.TICKER]",
|
| 88 |
+
description="Ticker de l'indice"
|
| 89 |
+
)
|
| 90 |
+
ref_unik_asset: str = Field(
|
| 91 |
+
...,
|
| 92 |
+
alias="PMVL[Ref Unik Asset]",
|
| 93 |
+
description="Référence unique de l'asset"
|
| 94 |
+
)
|
| 95 |
+
fund_code: str = Field(
|
| 96 |
+
...,
|
| 97 |
+
alias="PMVL[Selected Fund code]",
|
| 98 |
+
description="Code du fonds sélectionné"
|
| 99 |
+
)
|
| 100 |
+
col_3a: str = Field(
|
| 101 |
+
...,
|
| 102 |
+
alias="PMVL[3A]",
|
| 103 |
+
description="Classification 3A"
|
| 104 |
+
)
|
| 105 |
+
canton: str = Field(
|
| 106 |
+
...,
|
| 107 |
+
alias="PMVL[CANTON]",
|
| 108 |
+
description="Canton"
|
| 109 |
+
)
|
| 110 |
+
cic: str = Field(
|
| 111 |
+
...,
|
| 112 |
+
alias="PMVL[CIC]",
|
| 113 |
+
description="Classification CIC"
|
| 114 |
+
)
|
| 115 |
+
groupe: str = Field(
|
| 116 |
+
...,
|
| 117 |
+
alias="PMVL[GROUPE]",
|
| 118 |
+
description="Groupe"
|
| 119 |
+
)
|
| 120 |
+
ptf_name: str = Field(
|
| 121 |
+
...,
|
| 122 |
+
alias="PMVL[Ptf name]",
|
| 123 |
+
description="Nom du portefeuille"
|
| 124 |
+
)
|
| 125 |
+
|
| 126 |
+
class Config:
|
| 127 |
+
populate_by_name = True # Nouvelle syntaxe Pydantic V2
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
class PredictionResponse(BaseModel):
|
| 131 |
+
"""
|
| 132 |
+
Schéma de la réponse retournée par l'API.
|
| 133 |
+
"""
|
| 134 |
+
proba_bonne_estimation: float = Field(
|
| 135 |
+
...,
|
| 136 |
+
description="Probabilité (entre 0 et 1) que la PMVL soit une bonne estimation."
|
| 137 |
+
)
|
| 138 |
+
prediction: bool = Field(
|
| 139 |
+
...,
|
| 140 |
+
description="True si la PMVL est jugée bonne (proba >= seuil), False sinon."
|
| 141 |
+
)
|
| 142 |
+
seuil_applique: float = Field(
|
| 143 |
+
...,
|
| 144 |
+
description="Le seuil de probabilité utilisé pour cette décision (ex: 0.45)."
|
| 145 |
+
)
|
| 146 |
+
fund_code: str = Field(
|
| 147 |
+
...,
|
| 148 |
+
description="Rappel du code du fonds pour traçabilité."
|
| 149 |
+
)
|
| 150 |
+
ref_unik_asset: str = Field(
|
| 151 |
+
...,
|
| 152 |
+
description="Rappel de la référence de l'actif pour traçabilité."
|
| 153 |
+
)
|
models/pmvl_catboost_final.cbm
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:b81635e7410aee5b4d95784954ea3e6d9356bcfbfdce88f713aa9086bb31fc14
|
| 3 |
+
size 6085568
|
models/pmvl_feature_columns.txt
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
PMVL[ENTITE]
|
| 2 |
+
PMVL[ISIN]
|
| 3 |
+
PMVL[Orig. name]
|
| 4 |
+
PMVL[Parametres_Indices.TICKER]
|
| 5 |
+
PMVL[PFMIndice.Perf J / J-1]
|
| 6 |
+
PMVL[PRMP PMVL]
|
| 7 |
+
PMVL[PRMP VNC]
|
| 8 |
+
PMVL[PRMP MtM]
|
| 9 |
+
PMVL[Ref Unik Asset]
|
| 10 |
+
PMVL[Selected Fund code]
|
| 11 |
+
PMVL[3A]
|
| 12 |
+
PMVL[CANTON]
|
| 13 |
+
PMVL[CIC]
|
| 14 |
+
PMVL[GROUPE]
|
| 15 |
+
PMVL[Ptf name]
|
| 16 |
+
PMVL[Quantity]
|
| 17 |
+
PMVL[Purch. Val. (clean) (ptf cur.)]
|
| 18 |
+
PMVL[Quote Date]
|
| 19 |
+
PMVL[Quote]
|
| 20 |
+
PMVL[VNC Agrege dirty (ptf cur.)]
|
| 21 |
+
PMVL[PMVL Estimé]
|
| 22 |
+
position_group
|
requirements.txt
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
fastapi>=0.111.0
|
| 2 |
+
uvicorn[standard]>=0.30.1
|
| 3 |
+
pydantic>=2.7.0
|
| 4 |
+
catboost==1.2.10
|
| 5 |
+
pandas>=2.2.0
|
| 6 |
+
numpy>=1.26.0
|
| 7 |
+
pytest>=8.0.0
|
| 8 |
+
httpx>=0.27.0
|
tests/__init__.py
ADDED
|
File without changes
|
tests/test_api.py
ADDED
|
@@ -0,0 +1,85 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from fastapi.testclient import TestClient
|
| 2 |
+
from app.main import app
|
| 3 |
+
|
| 4 |
+
# Création du client de test FastAPI
|
| 5 |
+
client = TestClient(app)
|
| 6 |
+
|
| 7 |
+
# Payload valide (contenant tous les champs obligatoires définis dans schemas.py)
|
| 8 |
+
# On utilise ici les noms des variables Python (grâce à populate_by_name=True)
|
| 9 |
+
VALID_PAYLOAD = {
|
| 10 |
+
"holding_date": "2026-03-01",
|
| 11 |
+
"pmvl_estim": 1500.50,
|
| 12 |
+
"quantity": 100.0,
|
| 13 |
+
"purch_val_clean": 45000.0,
|
| 14 |
+
"quote": 510.0,
|
| 15 |
+
"vnc_agrege_dirty": 49000.0,
|
| 16 |
+
"entite": "ENTITE_TEST",
|
| 17 |
+
"isin": "FR0000000001",
|
| 18 |
+
"orig_name": "Asset Name Test",
|
| 19 |
+
"ticker": "TICKER_TEST",
|
| 20 |
+
"ref_unik_asset": "REF_12345",
|
| 21 |
+
"fund_code": "FUND_001",
|
| 22 |
+
"col_3a": "3A_TEST",
|
| 23 |
+
"canton": "CANTON_TEST",
|
| 24 |
+
"cic": "CIC_TEST",
|
| 25 |
+
"groupe": "GROUPE_TEST",
|
| 26 |
+
"ptf_name": "PTF_TEST"
|
| 27 |
+
}
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def test_health_check():
|
| 31 |
+
"""
|
| 32 |
+
Test 1 : Vérifie que le endpoint de diagnostic répond bien 200 OK.
|
| 33 |
+
"""
|
| 34 |
+
response = client.get("/health")
|
| 35 |
+
assert response.status_code == 200
|
| 36 |
+
assert response.json()["status"] == "ok"
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def test_predict_valid_data():
|
| 40 |
+
"""
|
| 41 |
+
Test 2 : Vérifie qu'un payload complet et valide renvoie bien une prédiction.
|
| 42 |
+
"""
|
| 43 |
+
response = client.post("/predict", json=VALID_PAYLOAD)
|
| 44 |
+
|
| 45 |
+
# Ligne ajoutée pour debug
|
| 46 |
+
print("RESPONSE STATUS:", response.status_code)
|
| 47 |
+
print("RESPONSE BODY:", response.json())
|
| 48 |
+
|
| 49 |
+
assert response.status_code == 200
|
| 50 |
+
|
| 51 |
+
# Vérifie la structure de la réponse
|
| 52 |
+
data = response.json()
|
| 53 |
+
assert "proba_bonne_estimation" in data
|
| 54 |
+
assert "prediction" in data
|
| 55 |
+
assert "seuil_applique" in data
|
| 56 |
+
assert data["fund_code"] == "FUND_001"
|
| 57 |
+
|
| 58 |
+
# La probabilité doit être entre 0 et 1
|
| 59 |
+
assert 0.0 <= data["proba_bonne_estimation"] <= 1.0
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
def test_predict_missing_required_field():
|
| 63 |
+
"""
|
| 64 |
+
Test 3 : Vérifie que l'API renvoie une erreur 422 si un champ obligatoire (ex: quantity) est manquant.
|
| 65 |
+
"""
|
| 66 |
+
invalid_payload = VALID_PAYLOAD.copy()
|
| 67 |
+
del invalid_payload["quantity"] # On supprime un champ obligatoire
|
| 68 |
+
|
| 69 |
+
response = client.post("/predict", json=invalid_payload)
|
| 70 |
+
|
| 71 |
+
# 422 Unprocessable Entity est le code standard de FastAPI pour une erreur de validation
|
| 72 |
+
assert response.status_code == 422
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
def test_predict_invalid_data_type():
|
| 76 |
+
"""
|
| 77 |
+
Test 4 : Vérifie que l'API renvoie une erreur 422 si un type de donnée est incorrect
|
| 78 |
+
(ex: une chaîne de caractères au lieu d'un float pour 'quote').
|
| 79 |
+
"""
|
| 80 |
+
invalid_payload = VALID_PAYLOAD.copy()
|
| 81 |
+
invalid_payload["quote"] = "Ceci_n_est_pas_un_chiffre"
|
| 82 |
+
|
| 83 |
+
response = client.post("/predict", json=invalid_payload)
|
| 84 |
+
|
| 85 |
+
assert response.status_code == 422
|