Compare commits

..
Author SHA1 Message Date
Johan LEROY b433e01fa8 fix(backend): départage aussi les égalités de timestamp dans latest_by_site
Backend / Lint, typage et tests (push) Successful in 1m18s
`latest_by_site` portait le même défaut que `latest_for_site` : `DISTINCT ON (site_id)`
ordonné sur `site_id, timestamp DESC` sans départage, alors que `uq_reading_source`
autorise deux lignes au même `site_id`+`timestamp` quand la `source` diffère.
`/stats/summary` pouvait donc afficher une consommation différente d'un appel à
l'autre pour un site alimenté par un backfill CSV et une écriture live.

Test `integration` dédié, qui échoue sans le correctif.
2026-09-18 10:28:04 +02:00
Johan LEROY 5eb74aa64a fix(backend): traite la revue de phyri0s sur la PR #84
Tri non déterministe : `latest_for_site` départage désormais les égalités de
timestamp par `reading_id` décroissant, comme `list_history`. `uq_reading_source`
autorise deux lignes au même `site_id`+`timestamp` quand la `source` diffère, donc
le `LIMIT 1` pouvait renvoyer l'une ou l'autre d'un appel à l'autre.

Tests : trois tests `integration` sur `latest_for_site` (plus récente, égalité de
timestamp, isolation par site). Le test d'égalité échoue sans le correctif ci-dessus.

Duplication : `DataQuality` et le repli vers `critical` sortent dans
`app/services/data_quality.py`, partagé par `stats.py`, `site.py` et `sensor.py`,
qui en portaient trois copies indépendantes. Supprime au passage deux
`# type: ignore[assignment]`.
2026-09-18 10:28:04 +02:00
Johan LEROY d167b64188 Merge remote-tracking branch 'origin/dev' into feat/endpoint-sites-current
# Conflicts:
#	apps/backend/app/repositories/reading.py
#	apps/backend/openapi.json
#	docs/architecture/20-backend.md
2026-09-17 12:17:05 +02:00
Johan LEROY 2f97e4d434 fix(backend): corrige formatage ruff et typage mypy sur sites/current
CI en échec sur ruff format (ligne trop longue) et mypy (retour Any non
annoté, assignation Literal non étroite). Corrige sans changer le
comportement.
2026-09-16 15:27:05 +02:00
Johan LEROY 07ea8d21dc feat(backend): expose GET /api/v1/sites/{site_id}/current pour l'issue #29
Ajoute la dernière mesure d'un site (SiteService.current), en réutilisant
la vérification d'existence déjà en place pour GET /sites/{site_id} :
SiteService gagne une dépendance ReadingRepository, sur le modèle de
composition déjà utilisé par StatsService/SensorService. Un site connu
sans lecture rend 200 avec les champs de mesure à null et
data_quality="critical" ; seul un site_id absent rend 404.
2026-09-16 15:25:14 +02:00
37 changed files with 561 additions and 3152 deletions
-40
View File
@@ -1,40 +0,0 @@
version: 2
updates:
# Frontend — npm
- package-ecosystem: "npm"
directory: "/apps/frontend"
schedule:
interval: "weekly"
open-pull-requests-limit: 5
groups:
frontend-dependencies:
patterns:
- "*"
# Backend — uv (lit pyproject.toml / uv.lock)
- package-ecosystem: "uv"
directory: "/apps/backend"
schedule:
interval: "weekly"
open-pull-requests-limit: 5
groups:
backend-dependencies:
patterns:
- "*"
# Les workflows GitHub Actions eux-mêmes ont aussi des dépendances à jour
- package-ecosystem: "github-actions"
directory: "/"
schedule:
interval: "weekly"
# Si un Dockerfile existe pour le backend
- package-ecosystem: "docker"
directory: "/apps/backend"
schedule:
interval: "weekly"
- package-ecosystem: "docker"
directory: "/apps/frontend"
schedule:
interval: "weekly"
-59
View File
@@ -1,59 +0,0 @@
name: ML
# Piège : la version de Python vient de ml/.python-version, et doit rester en 3.14 (cf.
# .github/workflows/backend.yml, même contrainte).
on:
push:
paths:
- "ml/**"
- ".github/workflows/ml.yml"
pull_request:
paths:
- "ml/**"
- ".github/workflows/ml.yml"
permissions:
contents: read
concurrency:
group: ml-${{ github.ref }}
cancel-in-progress: true
jobs:
verification:
name: Lint, typage et tests
runs-on: ubuntu-latest
defaults:
run:
working-directory: ml
steps:
- name: Récupère le dépôt
uses: actions/checkout@v4
- name: Installe uv
uses: astral-sh/setup-uv@v5
with:
enable-cache: true
cache-dependency-glob: ml/uv.lock
- name: Installe l'interpréteur déclaré par .python-version
run: uv python install
- name: Synchronise les dépendances sans dévier du verrou
run: uv sync --all-groups --frozen
- name: Vérifie le formatage
run: uv run ruff format --check .
- name: Analyse statique
run: uv run ruff check --output-format=github .
- name: Typage
run: uv run mypy enervision_ml tests
# Aucun test ne touche PostgreSQL ni MLflow distant : tout tourne sur donnees
# synthetiques ou un magasin SQLite local jetable (cf. ml/tests/test_train.py).
- name: Tests
run: uv run pytest
-8
View File
@@ -58,14 +58,6 @@ data/raw/*
monitoring/grafana/data/
monitoring/prometheus/data/
# ML : jeu de donnees, modeles entraines et suivi MLflow local, tous generes/volumineux
ml/data/
ml/models/*
!ml/models/.gitkeep
ml/mlruns/
ml/mlartifacts/
ml/mlflow.db
# IDE et OS
.idea/
.vscode/
+3 -22
View File
@@ -1,17 +1,15 @@
BACKEND := apps/backend
FRONTEND := apps/frontend
ML := ml
.DEFAULT_GOAL := help
.PHONY: help install install-backend install-frontend install-ml dev dev-backend dev-frontend \
.PHONY: help install install-backend install-frontend dev dev-backend dev-frontend \
lint format typecheck test test-cov test-integration check \
openapi docker-build db-up db-down db-reset db-logs db-psql migrate bootstrap-admin \
ml-lint ml-typecheck ml-test ml-check ml-train
openapi docker-build db-up db-down db-reset db-logs db-psql migrate bootstrap-admin
help: ## Liste les cibles disponibles
@grep -E '^[a-zA-Z_-]+:.*?## .*$$' $(MAKEFILE_LIST) | awk 'BEGIN {FS = ":.*?## "}; {printf " \033[36m%-16s\033[0m %s\n", $$1, $$2}'
install: install-backend install-frontend install-ml ## Installe les dépendances backend, frontend et ML
install: install-backend install-frontend ## Installe les dépendances backend et frontend
install-backend: ## Installe les dépendances du backend
cd $(BACKEND) && uv sync --all-groups
@@ -19,9 +17,6 @@ install-backend: ## Installe les dépendances du backend
install-frontend: ## Installe les dépendances du frontend
cd $(FRONTEND) && npm ci
install-ml: ## Installe les dépendances du pipeline ML
cd $(ML) && uv sync --all-groups
dev: ## Lance toute la stack (backend + frontend) en rechargement à chaud
@trap 'kill 0' EXIT INT TERM; \
$(MAKE) --no-print-directory dev-backend & \
@@ -60,20 +55,6 @@ check: lint typecheck test ## Chaîne de vérification complète
openapi: ## Régénère apps/backend/openapi.json depuis les routes déclarées
cd $(BACKEND) && uv run python -m app.cli export-openapi
ml-lint: ## Analyse statique du pipeline ML
cd $(ML) && uv run ruff check .
ml-typecheck: ## Vérifie le typage du pipeline ML
cd $(ML) && uv run mypy enervision_ml tests
ml-test: ## Exécute les tests du pipeline ML (donnees synthetiques, sans base ni serveur MLflow)
cd $(ML) && uv run pytest
ml-check: ml-lint ml-typecheck ml-test ## Chaîne de vérification complète du pipeline ML
ml-train: ## Entraine le modele LightGBM. CSV=chemin optionnel, sinon lit ML_DATABASE_URL
cd $(ML) && uv run python -m enervision_ml.train $(if $(CSV),--csv $(CSV),)
docker-build: ## Construit l'image du backend
docker build -t enervision-backend:local $(BACKEND)
-2
View File
@@ -25,7 +25,6 @@ Ce que la documentation apporte à chacun : [docs/architecture/00-vue-ensemble.m
| Infra | Terraform (k3s single-node) | `infra/terraform` | Initialise |
| CI/CD | GitHub Actions | `.github/workflows` | Backend en place |
| Monitoring | Prometheus, Grafana, Alertmanager | `monitoring` | A initialiser |
| ML | LightGBM, MLflow | `ml` | Entrainement initialise |
Le backend, la base et l'infrastructure (Terraform/k3s) sont initialises a ce stade. Le frontend
sert un tableau de bord sur `/dashboard`, dont les données proviennent de fixtures : les endpoints
@@ -54,7 +53,6 @@ L'etat detaille de chaque brique et les vues d'architecture sont dans
├── infra/terraform/
│ ├── modules/ Modules reutilisables
│ └── environments/ Racines Terraform, une par environnement
├── ml/ Pipeline d'entrainement LightGBM, suivi MLflow
├── monitoring/
│ ├── prometheus/ Collecte et regles d'alerte
│ ├── grafana/ Provisioning et dashboards
+1 -1
View File
@@ -142,7 +142,7 @@ UserServiceDep = Annotated[UserService, Depends(get_user_service)]
def get_site_service(session: SessionDep) -> SiteService:
return SiteService(sites=SiteRepository(session))
return SiteService(sites=SiteRepository(session), readings=ReadingRepository(session))
SiteServiceDep = Annotated[SiteService, Depends(get_site_service)]
+17 -1
View File
@@ -3,7 +3,7 @@ from fastapi import APIRouter, HTTPException, status
from app.api.deps import LecteurDep, SiteServiceDep
from app.api.openapi import REPONSE_VALIDATION, Reponses
from app.schemas.errors import ErrorResponse
from app.schemas.site import SiteResponse
from app.schemas.site import SiteCurrentResponse, SiteResponse
from app.services.site import SiteNotFoundError
router = APIRouter()
@@ -34,3 +34,19 @@ async def get_site(site_id: str, _: LecteurDep, service: SiteServiceDep) -> Site
status_code=status.HTTP_404_NOT_FOUND, detail="Site introuvable"
) from erreur
return SiteResponse.model_validate(site)
@router.get(
"/{site_id}/current",
response_model=SiteCurrentResponse,
summary="Dernière mesure d'un site",
responses=REPONSES_INTROUVABLE,
)
async def get_current(site_id: str, _: LecteurDep, service: SiteServiceDep) -> SiteCurrentResponse:
try:
actuel = await service.current(site_id)
except SiteNotFoundError as erreur:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, detail="Site introuvable"
) from erreur
return SiteCurrentResponse.model_validate(actuel)
+15 -2
View File
@@ -13,14 +13,27 @@ class ReadingRepository:
async def latest_by_site(self) -> Sequence[Reading]:
# `.distinct(site_id)` compile en `DISTINCT ON (site_id)` sous PostgreSQL : une seule
# ligne par site, la plus récente grâce à l'ordre composite qui suit.
# ligne par site, la plus récente grâce à l'ordre composite qui suit. `reading_id` départage
# les égalités de timestamp, que `uq_reading_source` autorise à `source` différente.
requete = (
select(Reading)
.distinct(Reading.site_id)
.order_by(Reading.site_id, Reading.timestamp.desc())
.order_by(Reading.site_id, Reading.timestamp.desc(), Reading.reading_id.desc())
)
return (await self._session.execute(requete)).scalars().all()
async def latest_for_site(self, site_id: str) -> Reading | None:
# Piège : `uq_reading_source` autorise deux lignes au même `site_id`+`timestamp` quand la
# `source` diffère. Sans `reading_id` en départage, le `LIMIT 1` renverrait au hasard.
requete = (
select(Reading)
.where(Reading.site_id == site_id)
.order_by(Reading.timestamp.desc(), Reading.reading_id.desc())
.limit(1)
)
lecture: Reading | None = await self._session.scalar(requete)
return lecture
async def list_history(
self,
*,
+20
View File
@@ -1,3 +1,6 @@
from datetime import datetime
from typing import Literal
from pydantic import BaseModel, ConfigDict
@@ -10,3 +13,20 @@ class SiteResponse(BaseModel):
location: str | None
capacity_kw: float | None
status: str | None
class SiteCurrentResponse(BaseModel):
model_config = ConfigDict(from_attributes=True)
timestamp: datetime | None
site_id: str
site_type: str
consumption_kw: float | None
consumption_kwh: float | None
voltage_v: float | None
current_a: float | None
power_factor: float | None
temperature_celsius: float | None
humidity_percent: float | None
null_reasons: list[str]
data_quality: Literal["good", "partial", "degraded", "critical"]
+18
View File
@@ -0,0 +1,18 @@
# Contrainte : `ck_reading_quality` accepte NULL et quatre valeurs seulement, alors que le contrat
# frontend n'a aucune valeur pour l'absence de qualité. `qualite_ou_critique()` replie donc sur
# `critical`, la seule des quatre qui n'induise pas une confiance qu'on n'a pas. `QUALITES_CONNUES`
# reste exposé pour les appelants qui doivent distinguer un `critical` stocké d'un repli.
from typing import Literal, get_args
DataQuality = Literal["good", "partial", "degraded", "critical"]
QUALITES_CONNUES: frozenset[str] = frozenset(get_args(DataQuality))
_PAR_VALEUR: dict[str, DataQuality] = {valeur: valeur for valeur in get_args(DataQuality)}
def qualite_ou_critique(valeur: str | None) -> DataQuality:
if valeur is None:
return "critical"
return _PAR_VALEUR.get(valeur, "critical")
+2 -3
View File
@@ -5,12 +5,11 @@ from typing import Literal
from app.models.energy import Reading, Site
from app.repositories.reading import ReadingRepository
from app.repositories.site import SiteRepository
from app.services.data_quality import qualite_ou_critique
CapteurStatus = Literal["ok", "failing"]
OverallStatus = Literal["ok", "degraded", "critical"]
QUALITES_CONNUES: frozenset[str] = frozenset({"good", "partial", "degraded", "critical"})
RAISON_VERS_CAPTEUR: dict[str, str] = {
"consumption_sensor_failure": "consumption",
"electrical_sensor_failure": "electrical",
@@ -80,7 +79,7 @@ def _sante_site(site: Site, derniere: Reading | None) -> SanteSite:
overall="critical",
)
qualite = derniere.data_quality if derniere.data_quality in QUALITES_CONNUES else "critical"
qualite = qualite_ou_critique(derniere.data_quality)
overall = _overall_depuis_qualite(qualite)
if overall == "critical":
+57 -1
View File
@@ -1,7 +1,11 @@
from collections.abc import Sequence
from dataclasses import dataclass
from datetime import datetime
from app.models.energy import Site
from app.repositories.reading import ReadingRepository
from app.repositories.site import SiteRepository
from app.services.data_quality import DataQuality, qualite_ou_critique
class SiteError(Exception):
@@ -12,9 +16,26 @@ class SiteNotFoundError(SiteError):
pass
@dataclass(frozen=True, slots=True)
class SiteCurrentReading:
timestamp: datetime | None
site_id: str
site_type: str
consumption_kw: float | None
consumption_kwh: float | None
voltage_v: float | None
current_a: float | None
power_factor: float | None
temperature_celsius: float | None
humidity_percent: float | None
null_reasons: list[str]
data_quality: DataQuality
class SiteService:
def __init__(self, *, sites: SiteRepository) -> None:
def __init__(self, *, sites: SiteRepository, readings: ReadingRepository) -> None:
self._sites = sites
self._readings = readings
async def list_all(self) -> Sequence[Site]:
return await self._sites.list_all()
@@ -24,3 +45,38 @@ class SiteService:
if site is None:
raise SiteNotFoundError(site_id)
return site
async def current(self, site_id: str) -> SiteCurrentReading:
site = await self.get_by_id(site_id)
derniere = await self._readings.latest_for_site(site_id)
if derniere is None:
return SiteCurrentReading(
timestamp=None,
site_id=site.site_id,
site_type=site.site_type,
consumption_kw=None,
consumption_kwh=None,
voltage_v=None,
current_a=None,
power_factor=None,
temperature_celsius=None,
humidity_percent=None,
null_reasons=[],
data_quality="critical",
)
return SiteCurrentReading(
timestamp=derniere.timestamp,
site_id=site.site_id,
site_type=site.site_type,
consumption_kw=derniere.consumption_kw,
consumption_kwh=derniere.consumption_kwh,
voltage_v=derniere.voltage_v,
current_a=derniere.current_a,
power_factor=derniere.power_factor,
temperature_celsius=derniere.temperature_celsius,
humidity_percent=derniere.humidity_percent,
null_reasons=derniere.null_reasons or [],
data_quality=qualite_ou_critique(derniere.data_quality),
)
+2 -9
View File
@@ -1,14 +1,10 @@
from dataclasses import dataclass
from datetime import UTC, datetime
from typing import Literal
from app.models.energy import Reading, Site
from app.repositories.reading import ReadingRepository
from app.repositories.site import SiteRepository
DataQuality = Literal["good", "partial", "degraded", "critical"]
QUALITES_CONNUES: frozenset[str] = frozenset({"good", "partial", "degraded", "critical"})
from app.services.data_quality import QUALITES_CONNUES, DataQuality, qualite_ou_critique
@dataclass(frozen=True, slots=True)
@@ -58,13 +54,10 @@ class StatsService:
@staticmethod
def _resume_site(site: Site, derniere: Reading | None) -> SiteConsumption:
capacite = site.capacity_kw or 0
# Piège : `data_quality` est nul dès qu'un site n'a jamais reçu de lecture, ou que le
# producteur n'a pas su la qualifier. Le contrat frontend n'a pas de valeur pour ce cas,
# `critical` est la seule des quatre qui n'induit pas une confiance qu'on n'a pas.
qualite: DataQuality = "critical"
consommation = None
if derniere is not None and derniere.data_quality in QUALITES_CONNUES:
qualite = derniere.data_quality # type: ignore[assignment]
qualite = qualite_ou_critique(derniere.data_quality)
consommation = derniere.consumption_kw
charge = (
+221
View File
@@ -921,6 +921,93 @@
}
}
},
"/api/v1/sites/{site_id}/current": {
"get": {
"tags": [
"sites"
],
"summary": "Dernière mesure d'un site",
"operationId": "get_current_api_v1_sites__site_id__current_get",
"security": [
{
"Jeton d'accès": []
}
],
"parameters": [
{
"name": "site_id",
"in": "path",
"required": true,
"schema": {
"type": "string",
"title": "Site Id"
}
}
],
"responses": {
"200": {
"description": "Successful Response",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/SiteCurrentResponse"
}
}
}
},
"500": {
"description": "Erreur interne. `correlation` identifie la trace côté serveur, qui n'est pas renvoyée au client.",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/InternalErrorResponse"
}
}
}
},
"401": {
"description": "Jeton absent, illisible, périmé, ou rendu caduc par un changement de rôle ou une désactivation. L'en-tête `WWW-Authenticate` porte la cause dans `error=`.",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
},
"403": {
"description": "Mot de passe provisoire à changer (`detail` vaut `password_change_required`).",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
},
"422": {
"description": "Corps invalide. Le détail nomme le champ fautif et le type d'erreur, jamais la valeur envoyée.",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ValidationErrorResponse"
}
}
}
},
"404": {
"description": "Aucun site ne porte cet identifiant.",
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/ErrorResponse"
}
}
}
}
}
}
},
"/api/v1/alerts": {
"get": {
"tags": [
@@ -2055,6 +2142,140 @@
],
"title": "SensorStatusResponse"
},
"SiteCurrentResponse": {
"properties": {
"timestamp": {
"anyOf": [
{
"type": "string",
"format": "date-time"
},
{
"type": "null"
}
],
"title": "Timestamp"
},
"site_id": {
"type": "string",
"title": "Site Id"
},
"site_type": {
"type": "string",
"title": "Site Type"
},
"consumption_kw": {
"anyOf": [
{
"type": "number"
},
{
"type": "null"
}
],
"title": "Consumption Kw"
},
"consumption_kwh": {
"anyOf": [
{
"type": "number"
},
{
"type": "null"
}
],
"title": "Consumption Kwh"
},
"voltage_v": {
"anyOf": [
{
"type": "number"
},
{
"type": "null"
}
],
"title": "Voltage V"
},
"current_a": {
"anyOf": [
{
"type": "number"
},
{
"type": "null"
}
],
"title": "Current A"
},
"power_factor": {
"anyOf": [
{
"type": "number"
},
{
"type": "null"
}
],
"title": "Power Factor"
},
"temperature_celsius": {
"anyOf": [
{
"type": "number"
},
{
"type": "null"
}
],
"title": "Temperature Celsius"
},
"humidity_percent": {
"anyOf": [
{
"type": "number"
},
{
"type": "null"
}
],
"title": "Humidity Percent"
},
"null_reasons": {
"items": {
"type": "string"
},
"type": "array",
"title": "Null Reasons"
},
"data_quality": {
"type": "string",
"enum": [
"good",
"partial",
"degraded",
"critical"
],
"title": "Data Quality"
}
},
"type": "object",
"required": [
"timestamp",
"site_id",
"site_type",
"consumption_kw",
"consumption_kwh",
"voltage_v",
"current_a",
"power_factor",
"temperature_celsius",
"humidity_percent",
"null_reasons",
"data_quality"
],
"title": "SiteCurrentResponse"
},
"SiteResponse": {
"properties": {
"site_id": {
+1
View File
@@ -31,6 +31,7 @@ ROUTES_A_ROLE = {
("POST", "/api/v1/users/{id}/password-reset"),
("GET", "/api/v1/sites"),
("GET", "/api/v1/sites/{site_id}"),
("GET", "/api/v1/sites/{site_id}/current"),
("GET", "/api/v1/alerts"),
("GET", "/api/v1/recommendations"),
("GET", "/api/v1/recommendations/{recommendation_id}"),
+51 -1
View File
@@ -1,4 +1,5 @@
from collections.abc import Callable, Iterator
from datetime import UTC, datetime
from uuid import uuid4
import pytest
@@ -9,7 +10,9 @@ from app.api.deps import get_current_principal, get_site_service
from app.core.principal import Principal
from app.core.roles import AccountKind, Role
from app.models.energy import Site
from app.services.site import SiteNotFoundError
from app.services.site import SiteCurrentReading, SiteNotFoundError
TIMESTAMP = datetime(2026, 9, 16, 12, 0, tzinfo=UTC)
def principal(role: Role = Role.LECTEUR) -> Principal:
@@ -33,10 +36,28 @@ def site(site_id: str = "site-1") -> Site:
)
def lecture_actuelle(site_id: str = "site-1") -> SiteCurrentReading:
return SiteCurrentReading(
timestamp=TIMESTAMP,
site_id=site_id,
site_type="industriel",
consumption_kw=87.34,
consumption_kwh=87.34,
voltage_v=401.2,
current_a=132.5,
power_factor=0.923,
temperature_celsius=22.1,
humidity_percent=58.4,
null_reasons=[],
data_quality="good",
)
class FauxService:
def __init__(self, erreur: Exception | None = None) -> None:
self._erreur = erreur
self.site = site()
self.actuel = lecture_actuelle()
async def list_all(self) -> list[Site]:
return [self.site]
@@ -46,6 +67,11 @@ class FauxService:
raise self._erreur
return self.site
async def current(self, site_id: str) -> SiteCurrentReading:
if self._erreur is not None:
raise self._erreur
return self.actuel
@pytest.fixture
def lecteur_connecte(app: FastAPI) -> Iterator[None]:
@@ -109,6 +135,30 @@ async def test_get_site_returns_404_for_an_unknown_site(
assert response.status_code == 404
async def test_get_current_returns_the_latest_reading(
servi: Callable[..., FauxService], client: AsyncClient
) -> None:
servi()
response = await client.get("/api/v1/sites/site-1/current")
assert response.status_code == 200
corps = response.json()
assert corps["site_id"] == "site-1"
assert corps["data_quality"] == "good"
assert corps["consumption_kw"] == 87.34
async def test_get_current_returns_404_for_an_unknown_site(
servi: Callable[..., FauxService], client: AsyncClient
) -> None:
servi(SiteNotFoundError("site-inconnu"))
response = await client.get("/api/v1/sites/site-inconnu/current")
assert response.status_code == 404
async def test_list_sites_reaches_the_repository_through_the_session(
lecteur_connecte: None, fake_session: Callable[..., None], client: AsyncClient
) -> None:
@@ -88,6 +88,73 @@ async def test_latest_by_site_returns_one_row_per_site(session: AsyncSession) ->
assert identifiants == {premier, second}
async def test_latest_by_site_breaks_a_timestamp_tie_on_the_last_written_reading(
session: AsyncSession,
) -> None:
site = await creer_site(session)
depot = ReadingRepository(session)
horodatage = datetime(2026, 9, 15, tzinfo=UTC)
await creer_lecture(
session, site_id=site.site_id, timestamp=horodatage, source="api_history", consumption_kw=10
)
derniere = await creer_lecture(
session, site_id=site.site_id, timestamp=horodatage, source="api_current", consumption_kw=42
)
resultats = await depot.latest_by_site()
retenues = [r.reading_id for r in resultats if r.site_id == site.site_id]
await session.rollback()
assert retenues == [derniere.reading_id]
async def test_latest_for_site_returns_the_most_recent_reading(session: AsyncSession) -> None:
site = await creer_site(session)
depot = ReadingRepository(session)
await creer_lecture(session, site_id=site.site_id, timestamp=datetime(2026, 9, 1, tzinfo=UTC))
recente = await creer_lecture(
session, site_id=site.site_id, timestamp=datetime(2026, 9, 15, tzinfo=UTC)
)
trouvee = await depot.latest_for_site(site.site_id)
reading_id = trouvee.reading_id if trouvee else None
await session.rollback()
assert reading_id == recente.reading_id
async def test_latest_for_site_breaks_a_timestamp_tie_on_the_last_written_reading(
session: AsyncSession,
) -> None:
site = await creer_site(session)
depot = ReadingRepository(session)
horodatage = datetime(2026, 9, 15, tzinfo=UTC)
await creer_lecture(session, site_id=site.site_id, timestamp=horodatage, source="api_history")
derniere = await creer_lecture(
session, site_id=site.site_id, timestamp=horodatage, source="api_current"
)
trouvee = await depot.latest_for_site(site.site_id)
reading_id = trouvee.reading_id if trouvee else None
await session.rollback()
assert reading_id == derniere.reading_id
async def test_latest_for_site_ignores_the_readings_of_the_other_sites(
session: AsyncSession,
) -> None:
sans_lecture = await creer_site(session)
autre = await creer_site(session)
depot = ReadingRepository(session)
await creer_lecture(session, site_id=autre.site_id)
trouvee = await depot.latest_for_site(sans_lecture.site_id)
await session.rollback()
assert trouvee is None
async def test_list_history_orders_the_readings_by_timestamp_descending(
session: AsyncSession,
) -> None:
+80 -7
View File
@@ -1,8 +1,13 @@
from dataclasses import dataclass, field
from datetime import UTC, datetime
import pytest
from app.models.energy import Site
from app.services.site import SiteNotFoundError, SiteService
TIMESTAMP = datetime(2026, 9, 16, 12, 0, tzinfo=UTC)
def site(site_id: str = "site-1") -> Site:
return Site(
@@ -15,6 +20,21 @@ def site(site_id: str = "site-1") -> Site:
)
@dataclass
class FauxLecture:
site_id: str
timestamp: datetime = TIMESTAMP
consumption_kw: float | None = 87.34
consumption_kwh: float | None = 87.34
voltage_v: float | None = 401.2
current_a: float | None = 132.5
power_factor: float | None = 0.923
temperature_celsius: float | None = 22.1
humidity_percent: float | None = 58.4
null_reasons: list[str] | None = field(default_factory=list)
data_quality: str | None = "good"
class FakeRepository:
def __init__(self, sites: list[Site]) -> None:
self._sites = sites
@@ -26,24 +46,77 @@ class FakeRepository:
return next((s for s in self._sites if s.site_id == site_id), None)
async def test_list_all_returns_the_repository_sites() -> None:
service = SiteService(sites=FakeRepository([site("a"), site("b")]))
class FauxDepotLectures:
def __init__(self, lectures: dict[str, FauxLecture]) -> None:
self._lectures = lectures
sites = await service.list_all()
async def latest_for_site(self, site_id: str) -> FauxLecture | None:
return self._lectures.get(site_id)
def service(sites: list[Site], lectures: dict[str, FauxLecture] | None = None) -> SiteService:
return SiteService(
sites=FakeRepository(sites), # type: ignore[arg-type]
readings=FauxDepotLectures(lectures or {}), # type: ignore[arg-type]
)
async def test_list_all_returns_the_repository_sites() -> None:
svc = service([site("a"), site("b")])
sites = await svc.list_all()
assert [s.site_id for s in sites] == ["a", "b"]
async def test_get_by_id_returns_the_matching_site() -> None:
service = SiteService(sites=FakeRepository([site("a")]))
svc = service([site("a")])
trouve = await service.get_by_id("a")
trouve = await svc.get_by_id("a")
assert trouve.site_id == "a"
async def test_get_by_id_raises_when_the_site_is_unknown() -> None:
service = SiteService(sites=FakeRepository([]))
svc = service([])
with pytest.raises(SiteNotFoundError):
await service.get_by_id("inconnu")
await svc.get_by_id("inconnu")
async def test_current_raises_when_the_site_is_unknown() -> None:
svc = service([])
with pytest.raises(SiteNotFoundError):
await svc.current("inconnu")
async def test_current_returns_every_field_as_null_when_the_site_has_no_reading() -> None:
svc = service([site("a")])
actuel = await svc.current("a")
assert actuel.timestamp is None
assert actuel.consumption_kw is None
assert actuel.data_quality == "critical"
assert actuel.null_reasons == []
async def test_current_copies_every_field_from_the_latest_reading() -> None:
svc = service([site("a")], {"a": FauxLecture(site_id="a")})
actuel = await svc.current("a")
assert actuel.timestamp == TIMESTAMP
assert actuel.site_type == "industriel"
assert actuel.consumption_kw == 87.34
assert actuel.voltage_v == 401.2
assert actuel.data_quality == "good"
async def test_current_treats_an_unknown_data_quality_as_critical() -> None:
svc = service([site("a")], {"a": FauxLecture(site_id="a", data_quality=None)})
actuel = await svc.current("a")
assert actuel.data_quality == "critical"
-101
View File
@@ -1,101 +0,0 @@
# 0005 - Modèle de prédiction de consommation : LightGBM
- Statut : accepté
- Date : 2026-09-17
## Contexte
Le schéma `prediction` contraint déjà la forme de la solution (deux cibles de régression,
`consumption_kw` instantané et `consumption_kwh` sur `period_minutes`, un statut
`insufficient_data` à détecter explicitement), mais aucun modèle n'était choisi. Trois
contraintes non négociables cadrent le choix, discutées dans l'issue #89 :
1. **EC06** (grille de notation individuelle) exige un modèle **entraîné, versionné avec
MLflow**, exposé via un endpoint fonctionnel, avec **surveillance du drift** en production.
2. **Aucun GPU dédié** : l'infra tourne on-premise sur une VM à 4 CPU / 8 Gio RAM (ou
`Standard_B2s`/`B2ms` côté Azure, 2 vCPU max) — Azure Machine Learning est de toute façon
bloqué par la politique Azure du projet.
3. **Délai serré** : le jalon J3 arrive à échéance le lendemain de la décision, J4 concentre déjà
26 issues sur 4 jours. Un modèle long à mettre en œuvre retarde la chaîne complète (service de
scoring #37, moteur de recommandations #38, tests ML #44/#45, tous bloqués par ce choix).
Le jeu de données est déjà disponible (`all_sites_combined.csv`, fourni par le formateur) : 7
sites, 2 ans au pas horaire (~17 500 lignes/site), avec `temperature_celsius`,
`humidity_percent`, `solar_irradiance_wm2` en régresseurs exogènes et des features calendaires
déjà dérivées.
## Options comparées
| Critère | Prophet | LightGBM/XGBoost | NeuralProphet | SARIMA | Holt-Winters | Mistral (LLM) |
|---|---|---|---|---|---|---|
| Saisonnalités multiples (jour/semaine/an) | Oui, nativement | Oui, via features engineered | Oui, nativement, + autorégression | Une seule, lourd à régler (SARIMAX) | Une seule, aucune | Non conçu pour ça |
| Régresseurs exogènes | Oui, mais doivent être connus dans le futur au moment de la prédiction | Oui, via lags/moyennes glissantes sur le passé | Oui, natif | Difficile en multivarié | Aucun support | Contexte de prompt seulement, non appris |
| Coût de calcul (VM sans GPU) | Faible | Faible | Élevé (deep learning) | Faible | Faible | Élevé à prohibitif |
| Versionnable MLflow | Oui, nativement | Oui, nativement | Pas de support direct | Oui, générique | Pas de support direct | Rien à versionner (pas un modèle entraîné) |
| Granularité | Un modèle par site (ou par site × métrique) | Un seul modèle global sur tous les sites | Un par site | Un par site | Un par site | — |
| Effort avant l'échéance | Faible | Moyen (feature engineering) | Élevé | Moyen à élevé | Faible en soi | Élevé, ou factice |
## Décision
**LightGBM, un seul modèle global** couvrant tous les sites, plutôt qu'un modèle par site
(Prophet) ou par famille de site. Cible : `consumption_kwh`, avec `period_minutes` comme feature
d'entrée plutôt que comme étape d'agrégation post-prédiction. Suivi et versioning via **MLflow**
(tracking + registre de modèles), sur le magasin local par défaut dans un premier temps —
l'hébergement sur l'infra k3s reste une question ouverte, non bloquante pour démarrer.
Raisons retenues, au-delà du tableau ci-dessus :
- **Un modèle global plutôt qu'un modèle par site** évite la fragilité des sites les moins
fournis en historique : ils bénéficient de ce qu'apprennent les autres sites, ce qu'un Prophet
par site ne permet pas.
- **Aucune dépendance à une prévision météo future.** Prophet exige que ses régresseurs
(`add_regressor`) soient connus au moment prédit ; `temperature_celsius`,
`humidity_percent` et `solar_irradiance_wm2` sont des mesures passées, pas des prévisions, et
aucune source de prévision météo n'existe dans le projet. LightGBM s'en sort avec des features
de lag/moyenne glissante calculées sur l'historique déjà présent dans `reading`, cf.
`ml/enervision_ml/features.py` — un choix qui vaut aussi bien à l'entraînement qu'au futur
scoring.
- **Apprentissage direct sur `consumption_kwh`** avec `period_minutes` en feature, sans étape
d'agrégation intermédiaire que la sortie continue de Prophet aurait demandée.
- **Coût de calcul compatible avec l'infra on-premise sans GPU.**
Débat complet, comparatif détaillé et décision finale : issue #89 (Johan, phyri0s,
ValentinDeFaria), actée en réunion d'équipe du 2026-09-17 et validée par l'ensemble de l'équipe.
## Conséquences
- Le pipeline d'entraînement (`ml/`, ce commit) lit `reading` + `site` par connexion PostgreSQL
directe et construit ses features par lags/moyennes glissantes plutôt que par régresseurs
contemporains, cf. `docs/ML-START.md`.
- Le rôle PostgreSQL dédié `enervision_ml` (lecture seule sur `reading`/`site`) n'est pas encore
provisionné : dette déjà assumée par l'ADR 0003 pour les comptes ETL/ML, `ML_DATABASE_URL`
pointe pour l'instant vers la même base que le backend applicatif en développement.
- Le service de scoring (#37), le moteur de recommandations (#38) et les tests de dérive
(#44/#45) restent à construire ; ils consommeront le même module `enervision_ml.features`, qui
doit rester strictement identique entre entraînement et scoring pour éviter un train/serve skew
silencieux.
- La surveillance de drift exigée par EC06 n'est pas encore implémentée : ce ticket ne livre que
l'entraînement et son suivi MLflow (paramètres, métriques, artefact modèle), pas le monitoring
en production.
- L'hébergement de MLflow sur l'infra k3s reste une question ouverte ; le magasin SQLite local
(`ml/mlflow.db`, ignoré par git) suffit pour l'instant à comparer des runs sur un poste.
## Alternatives écartées
- **Prophet** : proposition initiale, écartée après débat pour les raisons ci-dessus (modèle par
site, dépendance à une météo future indisponible, agrégation kWh en post-traitement). Reste un
candidat solide si un jour le projet doit produire une décomposition tendance/saisonnalité
explicable pour un usage différent.
- **Mistral (LLM)** : aucun produit dédié aux séries temporelles ; interroger un LLM généraliste
ne constitue pas un modèle entraîné et versionnable au sens MLflow, et le fine-tuning est hors
budget de calcul et hors délai.
- **SARIMA** : ne gère pas nativement plusieurs régresseurs exogènes ; réglage (p,d,q,P,D,Q) plus
long que le délai disponible.
- **NeuralProphet** : fait tout ce que fait Prophet et apprend en plus des motifs autorégressifs,
mais coûte plus cher en calcul (pas de GPU disponible) et n'a pas d'outil MLflow direct — piste
d'évolution possible, non engageante à ce stade.
- **Holt-Winters** : écarté d'entrée, pas seulement différé — aucun support de régresseurs
exogènes, alors que la météo et l'irradiance sont nécessaires ici.
- **CatBoost** : même famille que LightGBM, gère nativement les colonnes catégorielles (comme
`site_type`) sans encodage manuel. Non rejeté, différé : candidat à comparer si LightGBM
plafonne en précision.
-1
View File
@@ -77,7 +77,6 @@ collecteur ne vient le lire.
| Backend | FastAPI, Python 3.14 | `apps/backend` | `En cours` | Factory, configuration, journalisation, 2 sondes de santé, `/metrics`, contrat OpenAPI versionné, routes `sites`, `alerts`, `recommendations`, `stats/summary` et `readings` en lecture (endpoints → services → repositories → models) |
| Frontend | Angular 22, Node 24 | `apps/frontend` | `En cours` | Tableau de bord sur route `/dashboard`, deux services HTTP, graphiques Chart.js, données servies par des fixtures |
| Base | PostgreSQL 17 + TimescaleDB | `db` | `Fait` | Bootstrap de l'extension, base de test, chaîne Alembic. Schéma applicatif créé (`site`, `dataset`, `reading` en hypertable, `prediction`, `alert`, `recommendation`) |
| ML | LightGBM, MLflow | `ml` | `En cours` | Pipeline d'entraînement (features par lags/moyennes glissantes, baseline de persistance saisonnière, suivi MLflow local), voir [ADR 0005](../adr/0005-modele-prediction-lightgbm.md) et [ML-START.md](../../ML-START.md). Scoring, endpoint et surveillance de dérive pas encore construits |
| Infra | Terraform, k3s single-node | `infra/terraform` | `En cours` | Module d'installation du cluster. Jamais appliqué, aucune ressource Kubernetes déclarée |
| Monitoring | Prometheus, Grafana, Alertmanager | `monitoring` | `Cible` | Rien, hors le `/metrics` exposé par l'API |
| ETL | Apache Airflow | `etl/airflow` | `Cible` | Rien |
+6 -1
View File
@@ -142,6 +142,7 @@ Deux fichiers d'environnement, deux usages : `.env` à la racine alimente `docke
| POST | `/api/v1/users/{id}/password-reset` | Réinitialise et ferme les sessions. `admin` | 401, 403, 404, 422, 500 |
| GET | `/api/v1/sites` | Liste les sites. `lecteur` | 401, 403, 500 |
| GET | `/api/v1/sites/{site_id}` | Décrit un site. `lecteur` | 401, 403, 404, 422, 500 |
| GET | `/api/v1/sites/{site_id}/current` | Dernière mesure d'un site. `lecteur` | 401, 403, 404, 422, 500 |
| GET | `/api/v1/alerts` | Liste les alertes, filtrable par `site_id` et `severity`. `lecteur` | 401, 403, 422, 500 |
| GET | `/api/v1/recommendations` | Liste les recommandations. `lecteur` | 401, 403, 500 |
| GET | `/api/v1/recommendations/{recommendation_id}` | Décrit une recommandation. `lecteur` | 401, 403, 404, 422, 500 |
@@ -171,7 +172,11 @@ et `GET /recommendations/{recommendation_id}` reprennent le même gabarit à la
elle remonte à un site par sa seule `alert_id`, `alert` n'étant pas encore exposée. `GET
/stats/summary` et `GET /sensors/status` agrègent chacune deux repositories (`SiteRepository`,
`ReadingRepository`) dans un service dédié plutôt que d'exposer une table : elles n'entrent donc
pas dans ce gabarit route-par-table. Le contrat détaillé pour le frontend est dans
pas dans ce gabarit route-par-table. `GET /sites/{site_id}/current` reste sur le gabarit `sites`,
mais `SiteService` gagne la même seconde dépendance (`ReadingRepository`) pour restituer la
dernière `Reading` du site : un site connu sans lecture rend `200` avec tous les champs de mesure
à `null` et `data_quality="critical"`, seul un `site_id` absent de la base rend `404`. Le contrat
détaillé pour le frontend est dans
[31-contrat-authentification.md](31-contrat-authentification.md).
`GET /readings` reprend le même gabarit mais s'en écarte sur un point : `reading` est l'hypertable,
-1
View File
@@ -1 +0,0 @@
3.14
-88
View File
@@ -1,88 +0,0 @@
# ML EnerVision
Pipeline d'entrainement du modele de prevision de consommation energetique. Contexte complet :
[ADR 0005](../docs/adr/0005-modele-prediction-lightgbm.md) (choix du modele) et
[ML-START.md](../ML-START.md) (mecanisme d'acces aux donnees).
| Element | Choix |
|--------------|-----------------------------------------------|
| Python | 3.14 |
| Gestionnaire | uv (`uv.lock` fait foi) |
| Modele | LightGBM (regression, un seul modele global) |
| Suivi | MLflow (parametres, metriques, artefact) |
| Lint/format | ruff |
| Typage | mypy en mode strict |
| Tests | pytest, donnees synthetiques uniquement |
Projet Python independant de `apps/backend` : le service FastAPI n'a aucune raison d'embarquer
LightGBM/MLflow en dependance de production juste pour un script d'entrainement lance a la main.
## Installation
```bash
uv sync --all-groups
```
## Donnees
Deux sources, qui produisent le meme schema en sortie de `enervision_ml.data` (voir le module
pour le detail) :
- **CSV** (`--csv`), chemin de demarrage : lit directement `ml/data/all_sites_combined.csv`, le
jeu de donnees fourni pour le jalon J3. Ce dossier est ignore par git (gros fichier, local a
chaque poste) : recuperer le CSV et `dataset_metadata.json` aupres de l'equipe et les placer
dans `ml/data/` avant d'entrainer sur cette source.
- **PostgreSQL** (par defaut, sans `--csv`) : connexion directe a `reading` + `site` via
`ML_DATABASE_URL`, le chemin cible decrit dans `ML-START.md`. Le role PostgreSQL dedie
`enervision_ml` (lecture seule) n'est pas encore provisionne (dette assumee, cf. ADR 0003 et
ADR 0005) ; en attendant, pointer `ML_DATABASE_URL` vers la meme base que le backend suffit en
developpement.
## Entrainement
```bash
uv run python -m enervision_ml.train --csv data/all_sites_combined.csv
# ou, une fois la base peuplee et ML_DATABASE_URL positionnee :
uv run python -m enervision_ml.train
```
Ecrit le modele entraine dans `models/lightgbm-consumption.txt` (`Booster.save_model()`, dossier
ignore par git) et journalise la run dans MLflow : parametres, MAE/RMSE/MAPE du modele **et** de
la baseline de persistance saisonniere (consommation de la meme heure, une semaine avant), et
l'artefact modele. Sans `MLFLOW_TRACKING_URI`, MLflow ecrit dans un magasin SQLite local
(`./mlflow.db`, ignore par git) : `uv run mlflow ui` pour le consulter.
`--test-fraction` (0.15 par defaut) fixe la part la plus recente de l'historique reservee a la
validation. La coupure est **chronologique**, jamais un tirage aleatoire de lignes : un tirage
aleatoire laisserait des lignes de validation "voir" des lignes d'entrainement via leurs
lags/moyennes glissantes, une fuite qui masquerait un surapprentissage.
## Commandes
```bash
uv run ruff check . # lint
uv run ruff format . # format
uv run mypy enervision_ml tests # typage strict
uv run pytest # tests
```
Depuis la racine du monorepo, via le `Makefile` : `make install-ml`, `make ml-lint`,
`make ml-typecheck`, `make ml-test`, `make ml-check`, `make ml-train` (`CSV=chemin` optionnel).
## Ou ecrire les tests
Aucun test ne touche PostgreSQL ni un serveur MLflow distant : `enervision_ml.data.load_from_csv`
et le chargement CSV de test suffisent a exercer `build_features` sur des donnees reelles ou
synthetiques, et `enervision_ml.train.train()` accepte un `tracking_uri` SQLite isole (`tmp_path`
pytest) pour un test de bout en bout sans effet de bord. `enervision_ml.data.load_from_database`
n'est pas encore couvert : il n'existe aucune base PostgreSQL a interroger en CI ni dans cet
environnement de developpement pour le moment.
## Piege a connaitre
`enervision_ml.features.build_features` est **le seul endroit** qui doit construire les features
du modele, a l'entrainement comme au futur scoring (service #37, pas encore construit). Si les
deux divergent meme legerement (une fenetre de moyenne glissante calculee differemment, par
exemple), le modele recoit en production des features qui ne ressemblent plus a ce qu'il a
appris, et ses predictions deviennent silencieusement mauvaises sans qu'aucune erreur ne se
declenche. Ne jamais reecrire cette logique ailleurs : importer `enervision_ml.features`.
View File
-16
View File
@@ -1,16 +0,0 @@
"""Baseline de persistance saisonniere, la barre a depasser pour justifier LightGBM.
Predit la consommation de l'heure cible par celle de la meme heure, une semaine avant
(`consumption_kwh_lag_168h`) : une consommation energetique horaire est dominee par le cycle
hebdomadaire (jours ouvres contre week-end), donc ce naif-la est deja un concurrent serieux.
"""
import pandas as pd
from enervision_ml.features import TARGET_COLUMN
SEASONAL_LAG_COLUMN = f"{TARGET_COLUMN}_lag_168h"
def seasonal_persistence_predictions(features: pd.DataFrame) -> pd.Series:
return features[SEASONAL_LAG_COLUMN]
-38
View File
@@ -1,38 +0,0 @@
"""Configuration minimale du pipeline, lue depuis l'environnement.
Pas de `BaseSettings` Pydantic ici : contrairement a `apps/backend`, ce n'est pas un service qui
tourne en continu mais un script CLI lance a la main (cf. `docs/ML-START.md`), donc pas de
surface de configuration a valider au demarrage d'un processus long.
"""
import os
# Piege : ce n'est pas `DATABASE_URL` (celui du backend applicatif, proprietaire du schema).
# `docs/ML-START.md` et l'ADR 0003 designent un role PostgreSQL dedie et restreint en lecture,
# `enervision_ml`, non encore provisionne (dette assumee). Reutiliser `DATABASE_URL` par defaut
# ferait tourner l'entrainement avec les droits d'ecriture complets de l'application, en
# silence.
ML_DATABASE_URL_ENV = "ML_DATABASE_URL"
MLFLOW_EXPERIMENT_NAME = "consumption-forecast"
MLFLOW_TRACKING_URI_ENV = "MLFLOW_TRACKING_URI"
def database_url() -> str:
valeur = os.environ.get(ML_DATABASE_URL_ENV)
if not valeur:
raise RuntimeError(
f"{ML_DATABASE_URL_ENV} n'est pas defini. Elle doit pointer vers un role "
"PostgreSQL en lecture seule sur `reading`/`site` (voir docs/ML-START.md)."
)
return valeur
def mlflow_tracking_uri() -> str | None:
"""`None` laisse MLflow choisir son magasin local par defaut.
Piege : ce n'est plus `./mlruns` en clair depuis MLflow 3 (magasin fichier "maintenance
mode", refuse une URI `file:` explicite sauf `MLFLOW_ALLOW_FILE_STORE=true`), mais une base
SQLite locale (`./mlflow.db`).
"""
return os.environ.get(MLFLOW_TRACKING_URI_ENV)
-68
View File
@@ -1,68 +0,0 @@
"""Chargement des donnees d'entrainement.
Deux chemins, qui doivent produire le meme schema de sortie (colonnes `site_id`, `timestamp`,
`consumption_kwh`, `temperature_celsius`, `humidity_percent`, `solar_irradiance_wm2`,
`is_working_hours`, `site_type`, `capacity_kw`), consomme ensuite par `enervision_ml.features` :
- `load_from_database` : le chemin cible decrit dans `docs/ML-START.md`, connexion PostgreSQL
directe (`reading` + `site`), pas par l'API. C'est celui qu'utilisera le pipeline en
production, une fois le role PostgreSQL dedie `enervision_ml` provisionne (dette assumee,
documentee dans `CLAUDE.md` et l'ADR 0003 : pour l'instant, la meme chaine de connexion que le
backend applicatif convient en developpement).
- `load_from_csv` : chemin de demarrage, tant que la base locale n'est pas peuplee. Lit
directement `ml/data/all_sites_combined.csv` (jeu de donnees fourni pour le jalon J3, cf.
issue #89), le meme fichier que celui consomme par
`apps/backend/app/etl/historical_import.py`. `capacity_kw` n'existe pas dans ce CSV : la
colonne est renvoyee a `NaN`, que LightGBM gere nativement comme valeur manquante.
"""
from pathlib import Path
import pandas as pd
from sqlalchemy import text
from sqlalchemy.engine import Connectable
OUTPUT_COLUMNS = [
"site_id",
"timestamp",
"consumption_kwh",
"temperature_celsius",
"humidity_percent",
"solar_irradiance_wm2",
"is_working_hours",
"site_type",
"capacity_kw",
]
_READING_QUERY = text(
"""
SELECT
r.site_id,
r.timestamp,
r.consumption_kwh,
r.temperature_celsius,
r.humidity_percent,
r.solar_irradiance_wm2,
r.is_working_hours,
s.site_type,
s.capacity_kw
FROM reading r
JOIN site s ON s.site_id = r.site_id
ORDER BY r.site_id, r.timestamp
"""
)
def load_from_database(connection: Connectable) -> pd.DataFrame:
"""Lit l'historique complet `reading` + `site` depuis PostgreSQL."""
frame = pd.read_sql(_READING_QUERY, connection)
return frame[OUTPUT_COLUMNS]
def load_from_csv(csv_path: Path) -> pd.DataFrame:
"""Lit le jeu de donnees CSV historique (chemin de demarrage, hors base)."""
frame = pd.read_csv(csv_path, parse_dates=["timestamp"])
frame["capacity_kw"] = float("nan")
frame["is_working_hours"] = frame["is_working_hours"].astype(bool)
return frame[OUTPUT_COLUMNS]
-129
View File
@@ -1,129 +0,0 @@
"""Construction des features pour le modele de consommation.
Module partage entre l'entrainement et le futur scoring (cf. `docs/ML-START.md`) : la fonction
qui construit les features doit rester strictement identique des deux cotes, sous peine de
"train/serve skew" silencieux (le modele recoit en production des features qui ne ressemblent
plus a ce qu'il a appris).
"""
from collections.abc import Sequence
import pandas as pd
# Cible de l'entrainement : consommation en kWh, jamais consumption_kw (absent des lectures
# historiques CSV, cf. `apps/backend/app/etl/historical_import.py`).
TARGET_COLUMN = "consumption_kwh"
# Decalages horaires utilises pour les lags et moyennes glissantes : une heure avant, un jour
# avant (meme heure), une semaine avant (meme heure, meme jour) - saisonnalites usuelles d'une
# consommation energetique horaire.
LAG_HOURS: Sequence[int] = (1, 24, 168)
ROLLING_WINDOWS_HOURS: Sequence[int] = (24, 168)
STATIC_FEATURE_COLUMNS: Sequence[str] = ("site_type", "capacity_kw")
CALENDAR_FEATURE_COLUMNS: Sequence[str] = (
"hour",
"day_of_week",
"month",
"is_weekend",
"is_working_hours",
)
WEATHER_COLUMNS: Sequence[str] = (
"temperature_celsius",
"humidity_percent",
"solar_irradiance_wm2",
)
def build_features(frame: pd.DataFrame) -> pd.DataFrame:
"""Construit la matrice de features a partir de lectures brutes triees par site.
`frame` doit porter au minimum : `site_id`, `timestamp`, `consumption_kwh`,
`is_working_hours`, les trois colonnes meteo, et les colonnes statiques de site
(`site_type`, `capacity_kw`). Une ligne par `(site_id, timestamp)`, sans doublon.
Piege : la meteo n'entre dans les features que decalee (lag/moyenne glissante), jamais a
l'instant cible. A l'entrainement comme au scoring, la meteo au moment predit n'est pas une
mesure mais une prevision que le projet n'a pas — l'utiliser telle quelle romprait le
contrat entre entrainement et usage reel (la feature ne serait tout simplement plus
disponible en production). Cf. debat d'architecture dans l'issue #89.
"""
travail = frame.sort_values(["site_id", "timestamp"]).reset_index(drop=True)
calendrier = _calendar_features(travail["timestamp"])
decalees = _lagged_features(travail)
features = pd.concat(
[
travail[["site_id", "timestamp"]],
travail[list(STATIC_FEATURE_COLUMNS)],
calendrier,
travail[["is_working_hours"]],
decalees,
travail[[TARGET_COLUMN]],
],
axis=1,
)
# `period_minutes` : resolution temporelle de la cible. Les lectures historiques sont toutes
# au pas horaire (cf. `dataset_metadata.json`, `frequency: "1h""), donc une constante pour
# l'instant. Exposee comme feature plutot que supposee implicitement, pour que le modele
# puisse un jour apprendre sur d'autres resolutions sans reentrainement de zero.
features["period_minutes"] = 60
return features
def feature_columns() -> list[str]:
"""Liste ordonnee des colonnes d'entree du modele (hors identifiants et cible)."""
lag_columns = [f"consumption_kwh_lag_{h}h" for h in LAG_HOURS]
rolling_columns = [
f"{colonne}_rolling_mean_{fenetre}h"
for colonne in (TARGET_COLUMN, *WEATHER_COLUMNS)
for fenetre in ROLLING_WINDOWS_HOURS
]
weather_lag_columns = [f"{colonne}_lag_1h" for colonne in WEATHER_COLUMNS]
return [
*STATIC_FEATURE_COLUMNS,
*CALENDAR_FEATURE_COLUMNS,
"period_minutes",
*lag_columns,
*rolling_columns,
*weather_lag_columns,
]
def _calendar_features(timestamps: pd.Series) -> pd.DataFrame:
instants = pd.to_datetime(timestamps)
return pd.DataFrame(
{
"hour": instants.dt.hour,
"day_of_week": instants.dt.dayofweek,
"month": instants.dt.month,
"is_weekend": instants.dt.dayofweek.isin([5, 6]).astype(int),
}
)
def _lagged_features(travail: pd.DataFrame) -> pd.DataFrame:
par_site = travail.groupby("site_id", sort=False)
colonnes: dict[str, pd.Series] = {}
for decalage in LAG_HOURS:
colonnes[f"{TARGET_COLUMN}_lag_{decalage}h"] = par_site[TARGET_COLUMN].shift(decalage)
for colonne in (TARGET_COLUMN, *WEATHER_COLUMNS):
decale = par_site[colonne].shift(1)
for fenetre in ROLLING_WINDOWS_HOURS:
colonnes[f"{colonne}_rolling_mean_{fenetre}h"] = decale.groupby(
travail["site_id"]
).transform(lambda serie, fenetre=fenetre: serie.rolling(fenetre, min_periods=1).mean())
for colonne in WEATHER_COLUMNS:
colonnes[f"{colonne}_lag_1h"] = par_site[colonne].shift(1)
return pd.DataFrame(colonnes, index=travail.index)
-24
View File
@@ -1,24 +0,0 @@
"""Metriques de regression partagees entre le modele et la baseline."""
import numpy as np
import pandas as pd
from sklearn.metrics import mean_absolute_error, root_mean_squared_error
def regression_metrics(y_true: pd.Series, y_pred: pd.Series) -> dict[str, float]:
"""MAE, RMSE et MAPE (en %), sur les paires non nulles des deux series."""
valides = y_true.notna() & y_pred.notna()
reel = y_true[valides]
predit = y_pred[valides]
# MAPE diverge a consommation nulle : les mesures a zero (site a l'arret) sont exclues de ce
# seul ratio, pas des autres metriques.
non_nul = reel != 0
mape = float(np.mean(np.abs((reel[non_nul] - predit[non_nul]) / reel[non_nul])) * 100)
return {
"mae": float(mean_absolute_error(reel, predit)),
"rmse": float(root_mean_squared_error(reel, predit)),
"mape": mape,
"n_observations": int(valides.sum()),
}
-245
View File
@@ -1,245 +0,0 @@
"""Entrainement du modele LightGBM de prevision de consommation energetique.
CLI autonome, sur le meme gabarit que `apps/backend/app/etl/historical_import.py`
(argparse, connexion directe a la base). Cf. `docs/ML-START.md`, section 1.
uv run python -m enervision_ml.train --csv ../ml/data/all_sites_combined.csv
uv run python -m enervision_ml.train # lit ML_DATABASE_URL
Le modele entraine est ecrit en fichier (`Booster.save_model()`) et suivi par MLflow (parametres,
metriques, artefact). La base ne stocke jamais le modele lui-meme, seulement une reference vers
lui (`prediction.model_reference`, pose par le futur service de scoring - hors perimetre ici).
"""
import argparse
from pathlib import Path
from typing import Any
import lightgbm as lgb
import mlflow
import mlflow.lightgbm
import pandas as pd
from sqlalchemy import create_engine
from enervision_ml import config
from enervision_ml.baseline import seasonal_persistence_predictions
from enervision_ml.data import load_from_csv, load_from_database
from enervision_ml.features import TARGET_COLUMN, build_features, feature_columns
from enervision_ml.metrics import regression_metrics
CATEGORICAL_FEATURES = ["site_type"]
LIGHTGBM_PARAMS: dict[str, Any] = {
"objective": "regression",
"metric": "mae",
"learning_rate": 0.05,
"num_leaves": 63,
"min_data_in_leaf": 50,
"feature_fraction": 0.8,
"bagging_fraction": 0.8,
"bagging_freq": 1,
"verbosity": -1,
}
NUM_BOOST_ROUND = 1000
EARLY_STOPPING_ROUNDS = 50
DEFAULT_TEST_FRACTION = 0.15
def load_raw_frame(csv_path: Path | None) -> pd.DataFrame:
"""Lit les lectures brutes, depuis le CSV de demarrage ou depuis PostgreSQL."""
if csv_path is not None:
return load_from_csv(csv_path)
engine = create_engine(config.database_url())
try:
return load_from_database(engine)
finally:
engine.dispose()
def chronological_split(
features: pd.DataFrame, test_fraction: float
) -> tuple[pd.DataFrame, pd.DataFrame]:
"""Coupe par date de coupure, jamais par tirage aleatoire de lignes.
Une coupure aleatoire laisserait des lignes d'apres la coupure "voir" des lignes d'avant via
leurs lags/moyennes glissantes, une fuite qui masquerait un surapprentissage a l'evaluation.
"""
coupure = features["timestamp"].quantile(1 - test_fraction)
entrainement = features[features["timestamp"] < coupure]
validation = features[features["timestamp"] >= coupure]
return entrainement, validation
def prepare_dataset(frame: pd.DataFrame, columns: list[str]) -> tuple[pd.DataFrame, pd.Series]:
typee = frame.copy()
typee["site_type"] = typee["site_type"].astype("category")
return typee[columns], typee[TARGET_COLUMN]
def train(
*,
csv_path: Path | None,
model_output: Path,
test_fraction: float = DEFAULT_TEST_FRACTION,
tracking_uri: str | None = None,
) -> tuple[dict[str, float], dict[str, float]]:
"""Execute le pipeline complet et rend (metriques du modele, metriques de la baseline)."""
raw = load_raw_frame(csv_path)
features = build_features(raw)
columns = feature_columns()
# Les premieres 168h par site n'ont pas de lag hebdomadaire complet : ni entrainables, ni
# comparables a la baseline saisonniere qui en depend.
utilisable = features.dropna(subset=[TARGET_COLUMN, f"{TARGET_COLUMN}_lag_168h"])
entrainement, validation = chronological_split(utilisable, test_fraction)
if entrainement.empty or validation.empty:
raise ValueError(
"Fenetre d'entrainement ou de validation vide : jeu de donnees trop court pour "
f"test_fraction={test_fraction}."
)
X_train, y_train = prepare_dataset(entrainement, columns)
X_valid, y_valid = prepare_dataset(validation, columns)
train_set = lgb.Dataset(
X_train,
label=y_train,
categorical_feature=CATEGORICAL_FEATURES,
free_raw_data=False,
)
valid_set = lgb.Dataset(
X_valid,
label=y_valid,
reference=train_set,
categorical_feature=CATEGORICAL_FEATURES,
free_raw_data=False,
)
booster = lgb.train(
LIGHTGBM_PARAMS,
train_set,
num_boost_round=NUM_BOOST_ROUND,
valid_sets=[valid_set],
callbacks=[
lgb.early_stopping(EARLY_STOPPING_ROUNDS, verbose=False),
lgb.log_evaluation(period=0),
],
)
predictions = pd.Series(
booster.predict(X_valid, num_iteration=booster.best_iteration),
index=X_valid.index,
)
model_metrics = regression_metrics(y_valid, predictions)
baseline_metrics = regression_metrics(y_valid, seasonal_persistence_predictions(validation))
model_output.parent.mkdir(parents=True, exist_ok=True)
booster.save_model(str(model_output))
_log_to_mlflow(
tracking_uri=tracking_uri,
booster=booster,
model_metrics=model_metrics,
baseline_metrics=baseline_metrics,
n_train=len(X_train),
n_valid=len(X_valid),
test_fraction=test_fraction,
model_output=model_output,
)
return model_metrics, baseline_metrics
def _log_to_mlflow(
*,
tracking_uri: str | None,
booster: lgb.Booster,
model_metrics: dict[str, float],
baseline_metrics: dict[str, float],
n_train: int,
n_valid: int,
test_fraction: float,
model_output: Path,
) -> None:
uri = tracking_uri or config.mlflow_tracking_uri()
if uri is not None:
mlflow.set_tracking_uri(uri)
mlflow.set_experiment(config.MLFLOW_EXPERIMENT_NAME)
with mlflow.start_run():
mlflow.log_params(
{
**LIGHTGBM_PARAMS,
"num_boost_round": booster.best_iteration or NUM_BOOST_ROUND,
"test_fraction": test_fraction,
"n_train": n_train,
"n_valid": n_valid,
}
)
mlflow.log_metrics({f"model_{cle}": valeur for cle, valeur in model_metrics.items()})
mlflow.log_metrics({f"baseline_{cle}": valeur for cle, valeur in baseline_metrics.items()})
mlflow.lightgbm.log_model(booster, name="model")
mlflow.log_artifact(str(model_output))
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Entrainement du modele LightGBM EnerVision")
parser.add_argument(
"--csv",
type=Path,
default=None,
help=(
"Chemin vers le CSV historique (chemin de demarrage). Omis, lit ML_DATABASE_URL "
"et se connecte directement a PostgreSQL (reading + site)."
),
)
parser.add_argument(
"--model-output",
type=Path,
default=Path("models/lightgbm-consumption.txt"),
help="Chemin d'ecriture du modele entraine. Defaut : models/lightgbm-consumption.txt.",
)
parser.add_argument(
"--test-fraction",
type=float,
default=DEFAULT_TEST_FRACTION,
help=(
"Part la plus recente de l'historique reservee a la validation. "
f"Defaut : {DEFAULT_TEST_FRACTION}."
),
)
parser.add_argument(
"--mlflow-tracking-uri",
default=None,
help="Surcharge MLFLOW_TRACKING_URI. Omis, magasin SQLite local (./mlflow.db).",
)
return parser.parse_args()
def main() -> None:
args = parse_args()
model_metrics, baseline_metrics = train(
csv_path=args.csv,
model_output=args.model_output,
test_fraction=args.test_fraction,
tracking_uri=args.mlflow_tracking_uri,
)
print("Modele LightGBM :", model_metrics)
print("Baseline saisonniere (t-168h) :", baseline_metrics)
if model_metrics["mae"] < baseline_metrics["mae"]:
gain = (1 - model_metrics["mae"] / baseline_metrics["mae"]) * 100
print(f"LightGBM bat la baseline de {gain:.1f}% de MAE.")
else:
print("LightGBM ne bat pas la baseline saisonniere sur ce decoupage.")
if __name__ == "__main__":
main()
View File
-79
View File
@@ -1,79 +0,0 @@
[project]
name = "enervision-ml"
version = "0.1.0"
description = "Pipeline d'entrainement et de scoring du modele de prediction EnerVision (LightGBM)"
requires-python = ">=3.14,<3.15"
dependencies = [
"pandas>=3.0.5",
"sqlalchemy>=2.0.52",
"psycopg[binary]>=3.2",
"lightgbm>=4.6",
"scikit-learn>=1.7",
"mlflow>=3.0",
]
[dependency-groups]
dev = [
"ruff>=0.16.7",
"mypy>=2.3.1",
"pytest>=9.1.1",
"pandas-stubs>=3.0.5.260914",
]
[build-system]
requires = ["hatchling>=1.32.0"]
build-backend = "hatchling.build"
[tool.hatch.build.targets.wheel]
packages = ["enervision_ml"]
[tool.ruff]
line-length = 100
target-version = "py314"
src = ["enervision_ml", "tests"]
[tool.ruff.lint]
select = [
"E", "W",
"F",
"I",
"N",
"UP",
"B",
"C4",
"SIM",
"TID",
"RUF",
"S",
"PT",
]
# N806 : `X`/`y` (donnees/cible) est la convention scikit-learn/LightGBM, pas une variable mal
# nommee.
ignore = ["B008", "N806"]
[tool.ruff.lint.per-file-ignores]
"tests/**/*.py" = ["S101"]
[tool.ruff.lint.isort]
known-first-party = ["enervision_ml"]
[tool.ruff.format]
quote-style = "double"
[tool.mypy]
python_version = "3.14"
strict = true
warn_unreachable = true
[[tool.mypy.overrides]]
module = ["tests.*"]
disallow_untyped_defs = false
[[tool.mypy.overrides]]
module = ["lightgbm.*", "mlflow.*", "sklearn.*"]
ignore_missing_imports = true
[tool.pytest.ini_options]
testpaths = ["tests"]
addopts = "-q --strict-markers -m 'not integration'"
markers = ["integration: requiert une base PostgreSQL joignable"]
-11
View File
@@ -1,11 +0,0 @@
import pandas as pd
from enervision_ml.baseline import SEASONAL_LAG_COLUMN, seasonal_persistence_predictions
def test_seasonal_persistence_predictions_returns_the_168h_lag_column() -> None:
features = pd.DataFrame({SEASONAL_LAG_COLUMN: [1.0, 2.0, 3.0], "autre_colonne": [9, 9, 9]})
predictions = seasonal_persistence_predictions(features)
assert predictions.tolist() == [1.0, 2.0, 3.0]
-96
View File
@@ -1,96 +0,0 @@
from datetime import UTC, datetime, timedelta
from typing import cast
import pandas as pd
from enervision_ml.features import TARGET_COLUMN, build_features, feature_columns
def make_site_reading(
site_id: str, *, heures: int, depart: datetime, valeur: float = 10.0
) -> pd.DataFrame:
instants = [depart + timedelta(hours=h) for h in range(heures)]
return pd.DataFrame(
{
"site_id": site_id,
"timestamp": instants,
TARGET_COLUMN: [valeur + h for h in range(heures)],
"temperature_celsius": [15.0] * heures,
"humidity_percent": [50.0] * heures,
"solar_irradiance_wm2": [0.0] * heures,
"is_working_hours": [True] * heures,
"site_type": "office",
"capacity_kw": 100.0,
}
)
def two_site_frame(heures: int = 200) -> pd.DataFrame:
depart = datetime(2026, 1, 1, tzinfo=UTC)
return pd.concat(
[
make_site_reading("site-a", heures=heures, depart=depart, valeur=10.0),
make_site_reading("site-b", heures=heures, depart=depart, valeur=1000.0),
],
ignore_index=True,
)
def test_build_features_returns_every_declared_feature_column() -> None:
features = build_features(two_site_frame())
manquantes = set(feature_columns()) - set(features.columns)
assert manquantes == set()
def test_build_features_sets_a_constant_period_minutes() -> None:
features = build_features(two_site_frame())
assert (features["period_minutes"] == 60).all()
def test_build_features_lag_1h_matches_the_previous_hour_of_the_same_site() -> None:
features = build_features(two_site_frame(heures=200))
site_a = features[features["site_id"] == "site-a"].reset_index(drop=True)
assert site_a.loc[10, f"{TARGET_COLUMN}_lag_1h"] == site_a.loc[9, TARGET_COLUMN]
def test_build_features_lag_168h_is_nan_before_a_full_week_of_history() -> None:
features = build_features(two_site_frame(heures=200))
site_a = features[features["site_id"] == "site-a"].reset_index(drop=True)
assert pd.isna(site_a.loc[100, f"{TARGET_COLUMN}_lag_168h"])
assert not pd.isna(site_a.loc[168, f"{TARGET_COLUMN}_lag_168h"])
def test_build_features_never_leaks_lags_across_sites() -> None:
# site-b demarre a 1000 : si un lag de site-a s'y glissait, la valeur sortirait de son
# echelle (10, 11, 12, ...).
features = build_features(two_site_frame(heures=200))
site_b = features[features["site_id"] == "site-b"].reset_index(drop=True)
assert cast(float, site_b.loc[5, f"{TARGET_COLUMN}_lag_1h"]) >= 1000.0
def test_build_features_rolling_mean_excludes_the_current_hour() -> None:
# Valeurs constantes sauf la derniere ligne : si la moyenne glissante incluait l'heure
# courante, la constante ne resterait pas stable jusqu'au bout.
depart = datetime(2026, 1, 1, tzinfo=UTC)
frame = make_site_reading("site-a", heures=200, depart=depart, valeur=10.0)
frame[TARGET_COLUMN] = 10.0
frame.loc[frame.index[-1], TARGET_COLUMN] = 10_000.0
features = build_features(frame).reset_index(drop=True)
assert features.loc[len(features) - 1, f"{TARGET_COLUMN}_rolling_mean_24h"] == 10.0
def test_build_features_computes_calendar_fields_from_the_timestamp() -> None:
depart = datetime(2026, 1, 3, 6, tzinfo=UTC) # un samedi, 6h
features = build_features(make_site_reading("site-a", heures=1, depart=depart))
assert features.loc[0, "hour"] == 6
assert features.loc[0, "day_of_week"] == 5
assert features.loc[0, "is_weekend"] == 1
-45
View File
@@ -1,45 +0,0 @@
import pandas as pd
import pytest
from enervision_ml.metrics import regression_metrics
def test_regression_metrics_computes_mae_and_rmse_on_known_values() -> None:
y_true = pd.Series([10.0, 20.0, 30.0])
y_pred = pd.Series([12.0, 18.0, 33.0])
resultat = regression_metrics(y_true, y_pred)
assert resultat["mae"] == pytest.approx(7 / 3)
assert resultat["n_observations"] == 3
def test_regression_metrics_ignores_rows_with_a_missing_value() -> None:
y_true = pd.Series([10.0, None, 30.0])
y_pred = pd.Series([12.0, 18.0, None])
resultat = regression_metrics(y_true, y_pred)
assert resultat["n_observations"] == 1
assert resultat["mae"] == 2.0
def test_regression_metrics_excludes_zero_actuals_from_mape_only() -> None:
y_true = pd.Series([0.0, 10.0])
y_pred = pd.Series([5.0, 12.0])
resultat = regression_metrics(y_true, y_pred)
assert resultat["n_observations"] == 2
assert resultat["mape"] == pytest.approx(20.0)
def test_metrics_are_zero_for_a_perfect_prediction() -> None:
y_true = pd.Series([10.0, 20.0])
y_pred = pd.Series([10.0, 20.0])
resultat = regression_metrics(y_true, y_pred)
assert resultat["mae"] == 0.0
assert resultat["rmse"] == 0.0
assert resultat["mape"] == 0.0
-76
View File
@@ -1,76 +0,0 @@
from datetime import UTC, datetime, timedelta
from pathlib import Path
import numpy as np
import pandas as pd
from enervision_ml.features import TARGET_COLUMN, build_features, feature_columns
from enervision_ml.train import chronological_split, prepare_dataset, train
def make_frame(site_id: str, *, heures: int, depart: datetime) -> pd.DataFrame:
instants = [depart + timedelta(hours=h) for h in range(heures)]
rng = np.random.default_rng(42)
return pd.DataFrame(
{
"site_id": site_id,
"timestamp": instants,
TARGET_COLUMN: 100.0 + 10.0 * np.sin(np.arange(heures) / 24) + rng.normal(0, 1, heures),
"temperature_celsius": 15.0,
"humidity_percent": 50.0,
"solar_irradiance_wm2": 0.0,
"is_working_hours": True,
"site_type": "office",
"capacity_kw": 100.0,
}
)
def test_chronological_split_puts_the_most_recent_rows_in_validation() -> None:
depart = datetime(2026, 1, 1, tzinfo=UTC)
features = make_frame("site-a", heures=200, depart=depart)
entrainement, validation = chronological_split(features, test_fraction=0.2)
assert entrainement["timestamp"].max() < validation["timestamp"].min()
# La coupure vient d'un quantile sur les dates : une approximation du taux demande, pas un
# decompte exact de lignes.
assert abs(len(validation) - 0.2 * len(features)) <= 2
def test_prepare_dataset_types_site_type_as_a_pandas_category() -> None:
depart = datetime(2026, 1, 1, tzinfo=UTC)
features = build_features(make_frame("site-a", heures=200, depart=depart))
X, y = prepare_dataset(features, feature_columns())
assert X["site_type"].dtype.name == "category"
assert y.name == TARGET_COLUMN
def test_train_runs_end_to_end_on_synthetic_data_and_beats_a_dummy_baseline(
tmp_path: Path,
) -> None:
depart = datetime(2026, 1, 1, tzinfo=UTC)
frame = pd.concat(
[
make_frame("site-a", heures=400, depart=depart),
make_frame("site-b", heures=400, depart=depart),
],
ignore_index=True,
)
csv_path = tmp_path / "synthetic.csv"
frame.to_csv(csv_path, index=False)
model_metrics, baseline_metrics = train(
csv_path=csv_path,
model_output=tmp_path / "model.txt",
test_fraction=0.2,
tracking_uri=f"sqlite:///{tmp_path / 'mlflow.db'}",
)
assert (tmp_path / "model.txt").exists()
assert model_metrics["n_observations"] > 0
assert model_metrics["mae"] >= 0
assert baseline_metrics["n_observations"] == model_metrics["n_observations"]
Generated
-1977
View File
File diff suppressed because it is too large Load Diff