From e9376a98bfd4354b71d56a0057f9600cbbce509c Mon Sep 17 00:00:00 2001 From: Dorian Date: Fri, 18 Sep 2026 11:06:04 +0200 Subject: [PATCH] feat(ml,backend): implemente le service de scoring et GET /predictions (#37) --- Makefile | 5 +- apps/backend/app/api/deps.py | 11 + apps/backend/app/api/openapi.py | 7 + .../app/api/v1/endpoints/predictions.py | 18 ++ apps/backend/app/api/v1/router.py | 4 + apps/backend/app/repositories/prediction.py | 28 ++ apps/backend/app/schemas/prediction.py | 43 +++ apps/backend/app/services/prediction.py | 68 +++++ apps/backend/openapi.json | 197 ++++++++++++ apps/backend/tests/api/test_openapi.py | 1 + apps/backend/tests/api/test_predictions.py | 82 +++++ .../tests/repositories/test_prediction.py | 92 ++++++ .../backend/tests/services/test_prediction.py | 121 ++++++++ docs/architecture/00-vue-ensemble.md | 4 +- docs/architecture/20-backend.md | 36 ++- ml/README.md | 56 +++- ml/enervision_ml/data.py | 64 +++- ml/enervision_ml/score.py | 281 ++++++++++++++++++ ml/pyproject.toml | 6 +- ml/tests/test_data.py | 34 +++ ml/tests/test_score.py | 230 ++++++++++++++ 21 files changed, 1369 insertions(+), 19 deletions(-) create mode 100644 apps/backend/app/api/v1/endpoints/predictions.py create mode 100644 apps/backend/app/repositories/prediction.py create mode 100644 apps/backend/app/schemas/prediction.py create mode 100644 apps/backend/app/services/prediction.py create mode 100644 apps/backend/tests/api/test_predictions.py create mode 100644 apps/backend/tests/repositories/test_prediction.py create mode 100644 apps/backend/tests/services/test_prediction.py create mode 100644 ml/enervision_ml/score.py create mode 100644 ml/tests/test_data.py create mode 100644 ml/tests/test_score.py diff --git a/Makefile b/Makefile index 0bb1dcb..2a4b3d2 100644 --- a/Makefile +++ b/Makefile @@ -6,7 +6,7 @@ ML := ml .PHONY: help install install-backend install-frontend install-ml 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 + ml-lint ml-typecheck ml-test ml-check ml-train ml-score 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}' @@ -74,6 +74,9 @@ ml-check: ml-lint ml-typecheck ml-test ## Chaîne de vérification complète du 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),) +ml-score: ## Score le prochain pas horaire et l'ecrit dans `prediction`. CSV=chemin optionnel + cd $(ML) && uv run python -m enervision_ml.score $(if $(CSV),--csv $(CSV),) + docker-build: ## Construit l'image du backend docker build -t enervision-backend:local $(BACKEND) diff --git a/apps/backend/app/api/deps.py b/apps/backend/app/api/deps.py index ae5df24..0948cc8 100644 --- a/apps/backend/app/api/deps.py +++ b/apps/backend/app/api/deps.py @@ -27,6 +27,7 @@ from app.repositories.audit_log import AuditLogRepository from app.repositories.login_attempt import LoginAttemptRepository from app.repositories.password_reset_attempt import PasswordResetAttemptRepository from app.repositories.password_reset_token import PasswordResetTokenRepository +from app.repositories.prediction import PredictionRepository from app.repositories.reading import ReadingRepository from app.repositories.recommendation import RecommendationRepository from app.repositories.refresh_token import RefreshTokenRepository @@ -34,6 +35,7 @@ from app.repositories.site import SiteRepository from app.repositories.user import UserRepository from app.services.alert import AlertService from app.services.auth import AuthService, LoginPolicy, PasswordResetPolicy +from app.services.prediction import PredictionService from app.services.reading import ReadingService from app.services.recommendation import RecommendationService from app.services.sensor import SensorService @@ -212,6 +214,15 @@ def get_sensor_service(session: SessionDep) -> SensorService: SensorServiceDep = Annotated[SensorService, Depends(get_sensor_service)] +def get_prediction_service(session: SessionDep) -> PredictionService: + return PredictionService( + sites=SiteRepository(session), predictions=PredictionRepository(session) + ) + + +PredictionServiceDep = Annotated[PredictionService, Depends(get_prediction_service)] + + async def get_current_principal( credentials: CredentialsDep, session: SessionDep, diff --git a/apps/backend/app/api/openapi.py b/apps/backend/app/api/openapi.py index 9937467..11b2604 100644 --- a/apps/backend/app/api/openapi.py +++ b/apps/backend/app/api/openapi.py @@ -83,6 +83,13 @@ TAGS: Final[list[dict[str, Any]]] = [ "name": "sensors", "description": "État de santé des capteurs par site. Réservé au rôle `admin`.", }, + { + "name": "predictions", + "description": ( + "Dernière prévision de consommation par site, calculée hors ligne par le pipeline " + "de scoring (`ml/`) et simplement lue ici. Accessible à partir du rôle `lecteur`." + ), + }, ] cookie_de_rafraichissement = APIKeyCookie( diff --git a/apps/backend/app/api/v1/endpoints/predictions.py b/apps/backend/app/api/v1/endpoints/predictions.py new file mode 100644 index 0000000..61a534a --- /dev/null +++ b/apps/backend/app/api/v1/endpoints/predictions.py @@ -0,0 +1,18 @@ +from fastapi import APIRouter + +from app.api.deps import LecteurDep, PredictionServiceDep +from app.schemas.prediction import PredictionSummaryResponse + +router = APIRouter() + + +@router.get( + "", + response_model=PredictionSummaryResponse, + summary="Dernière prédiction de consommation par site", +) +async def get_predictions( + _: LecteurDep, service: PredictionServiceDep +) -> PredictionSummaryResponse: + resume = await service.summary() + return PredictionSummaryResponse.model_validate(resume) diff --git a/apps/backend/app/api/v1/router.py b/apps/backend/app/api/v1/router.py index 8c0bd9e..6079acf 100644 --- a/apps/backend/app/api/v1/router.py +++ b/apps/backend/app/api/v1/router.py @@ -5,6 +5,7 @@ from app.api.v1.endpoints import ( alerts, auth, health, + predictions, readings, recommendations, sensors, @@ -34,3 +35,6 @@ api_router.include_router( api_router.include_router( sensors.router, prefix="/sensors", tags=["sensors"], responses=REPONSES_ADMIN ) +api_router.include_router( + predictions.router, prefix="/predictions", tags=["predictions"], responses=REPONSES_LECTEUR +) diff --git a/apps/backend/app/repositories/prediction.py b/apps/backend/app/repositories/prediction.py new file mode 100644 index 0000000..f79311a --- /dev/null +++ b/apps/backend/app/repositories/prediction.py @@ -0,0 +1,28 @@ +from collections.abc import Sequence + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from app.models.energy import Prediction + + +class PredictionRepository: + def __init__(self, session: AsyncSession) -> None: + self._session = session + + async def latest_by_site(self) -> Sequence[Prediction]: + # `.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. Même mécanisme que + # `ReadingRepository.latest_by_site`. Trié sur `target_at` (couvert par + # `ix_prediction_site_target`) plutôt que `created_at` : c'est la prévision la plus + # récente qui compte pour un tableau de bord, pas forcément le dernier run de scoring. + requete = ( + select(Prediction) + .distinct(Prediction.site_id) + .order_by( + Prediction.site_id, + Prediction.target_at.desc(), + Prediction.prediction_id.desc(), + ) + ) + return (await self._session.scalars(requete)).all() diff --git a/apps/backend/app/schemas/prediction.py b/apps/backend/app/schemas/prediction.py new file mode 100644 index 0000000..b7eeea2 --- /dev/null +++ b/apps/backend/app/schemas/prediction.py @@ -0,0 +1,43 @@ +from datetime import datetime +from enum import StrEnum + +from pydantic import BaseModel, ConfigDict + + +class PredictionTargetMetric(StrEnum): + CONSUMPTION_KWH = "consumption_kwh" + CONSUMPTION_KW = "consumption_kw" + + +class PredictionStatus(StrEnum): + AVAILABLE = "available" + INSUFFICIENT_DATA = "insufficient_data" + ERROR = "error" + + +class SitePredictionResponse(BaseModel): + model_config = ConfigDict(from_attributes=True) + + target_at: datetime + target_metric: PredictionTargetMetric + period_minutes: int | None + predicted_value: float | None + status: PredictionStatus + failure_reason: str | None + model_reference: str + created_at: datetime + + +class SitePredictionSummaryResponse(BaseModel): + model_config = ConfigDict(from_attributes=True) + + site_id: str + site_name: str + prediction: SitePredictionResponse | None + + +class PredictionSummaryResponse(BaseModel): + model_config = ConfigDict(from_attributes=True) + + timestamp: datetime + sites: list[SitePredictionSummaryResponse] diff --git a/apps/backend/app/services/prediction.py b/apps/backend/app/services/prediction.py new file mode 100644 index 0000000..6235bf2 --- /dev/null +++ b/apps/backend/app/services/prediction.py @@ -0,0 +1,68 @@ +from dataclasses import dataclass +from datetime import UTC, datetime + +from app.models.energy import Prediction, Site +from app.repositories.prediction import PredictionRepository +from app.repositories.site import SiteRepository + + +@dataclass(frozen=True, slots=True) +class SitePrediction: + target_at: datetime + target_metric: str + period_minutes: int | None + predicted_value: float | None + status: str + failure_reason: str | None + model_reference: str + created_at: datetime + + +@dataclass(frozen=True, slots=True) +class SitePredictionSummary: + site_id: str + site_name: str + prediction: SitePrediction | None + + +@dataclass(frozen=True, slots=True) +class PredictionSummary: + timestamp: datetime + sites: list[SitePredictionSummary] + + +class PredictionService: + def __init__(self, sites: SiteRepository, predictions: PredictionRepository) -> None: + self._sites = sites + self._predictions = predictions + + async def summary(self) -> PredictionSummary: + sites = await self._sites.list_all() + dernieres = {p.site_id: p for p in await self._predictions.latest_by_site()} + + return PredictionSummary( + timestamp=datetime.now(UTC), + sites=[_resume_site(site, dernieres.get(site.site_id)) for site in sites], + ) + + +def _resume_site(site: Site, derniere: Prediction | None) -> SitePredictionSummary: + # Piège : l'absence de ligne signifie « jamais scoré », pas une valeur pseudo-statut, qui + # n'existe pas dans la contrainte de la table. `prediction` reste `None` plutôt que de + # fabriquer un statut absent du domaine `available`/`insufficient_data`/`error`. + prediction = None + if derniere is not None: + prediction = SitePrediction( + target_at=derniere.target_at, + target_metric=derniere.target_metric, + period_minutes=derniere.period_minutes, + predicted_value=derniere.predicted_value, + status=derniere.status, + failure_reason=derniere.failure_reason, + model_reference=derniere.model_reference, + created_at=derniere.created_at, + ) + + return SitePredictionSummary( + site_id=site.site_id, site_name=site.site_name, prediction=prediction + ) diff --git a/apps/backend/openapi.json b/apps/backend/openapi.json index 25fc9ee..08b6a77 100644 --- a/apps/backend/openapi.json +++ b/apps/backend/openapi.json @@ -1632,6 +1632,62 @@ } ] } + }, + "/api/v1/predictions": { + "get": { + "tags": [ + "predictions" + ], + "summary": "Dernière prédiction de consommation par site", + "operationId": "get_predictions_api_v1_predictions_get", + "responses": { + "200": { + "description": "Successful Response", + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/PredictionSummaryResponse" + } + } + } + }, + "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" + } + } + } + } + }, + "security": [ + { + "Jeton d'accès": [] + } + ] + } } }, "components": { @@ -1885,6 +1941,45 @@ ], "title": "PasswordChangeRequest" }, + "PredictionStatus": { + "type": "string", + "enum": [ + "available", + "insufficient_data", + "error" + ], + "title": "PredictionStatus" + }, + "PredictionSummaryResponse": { + "properties": { + "timestamp": { + "type": "string", + "format": "date-time", + "title": "Timestamp" + }, + "sites": { + "items": { + "$ref": "#/components/schemas/SitePredictionSummaryResponse" + }, + "type": "array", + "title": "Sites" + } + }, + "type": "object", + "required": [ + "timestamp", + "sites" + ], + "title": "PredictionSummaryResponse" + }, + "PredictionTargetMetric": { + "type": "string", + "enum": [ + "consumption_kwh", + "consumption_kw" + ], + "title": "PredictionTargetMetric" + }, "PrincipalResponse": { "properties": { "id": { @@ -2297,6 +2392,104 @@ ], "title": "SensorStatusResponse" }, + "SitePredictionResponse": { + "properties": { + "target_at": { + "type": "string", + "format": "date-time", + "title": "Target At" + }, + "target_metric": { + "$ref": "#/components/schemas/PredictionTargetMetric" + }, + "period_minutes": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Period Minutes" + }, + "predicted_value": { + "anyOf": [ + { + "type": "number" + }, + { + "type": "null" + } + ], + "title": "Predicted Value" + }, + "status": { + "$ref": "#/components/schemas/PredictionStatus" + }, + "failure_reason": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Failure Reason" + }, + "model_reference": { + "type": "string", + "title": "Model Reference" + }, + "created_at": { + "type": "string", + "format": "date-time", + "title": "Created At" + } + }, + "type": "object", + "required": [ + "target_at", + "target_metric", + "period_minutes", + "predicted_value", + "status", + "failure_reason", + "model_reference", + "created_at" + ], + "title": "SitePredictionResponse" + }, + "SitePredictionSummaryResponse": { + "properties": { + "site_id": { + "type": "string", + "title": "Site Id" + }, + "site_name": { + "type": "string", + "title": "Site Name" + }, + "prediction": { + "anyOf": [ + { + "$ref": "#/components/schemas/SitePredictionResponse" + }, + { + "type": "null" + } + ] + } + }, + "type": "object", + "required": [ + "site_id", + "site_name", + "prediction" + ], + "title": "SitePredictionSummaryResponse" + }, "SiteResponse": { "properties": { "site_id": { @@ -2752,6 +2945,10 @@ { "name": "sensors", "description": "État de santé des capteurs par site. Réservé au rôle `admin`." + }, + { + "name": "predictions", + "description": "Dernière prévision de consommation par site, calculée hors ligne par le pipeline de scoring (`ml/`) et simplement lue ici. Accessible à partir du rôle `lecteur`." } ] } diff --git a/apps/backend/tests/api/test_openapi.py b/apps/backend/tests/api/test_openapi.py index aa5546b..32bcb21 100644 --- a/apps/backend/tests/api/test_openapi.py +++ b/apps/backend/tests/api/test_openapi.py @@ -37,6 +37,7 @@ ROUTES_A_ROLE = { ("GET", "/api/v1/stats/summary"), ("GET", "/api/v1/readings"), ("GET", "/api/v1/sensors/status"), + ("GET", "/api/v1/predictions"), } diff --git a/apps/backend/tests/api/test_predictions.py b/apps/backend/tests/api/test_predictions.py new file mode 100644 index 0000000..cbc1cb6 --- /dev/null +++ b/apps/backend/tests/api/test_predictions.py @@ -0,0 +1,82 @@ +from collections.abc import Callable, Iterator +from datetime import UTC, datetime +from uuid import uuid4 + +import pytest +from fastapi import FastAPI +from httpx import AsyncClient + +from app.api.deps import get_current_principal, get_prediction_service +from app.core.principal import Principal +from app.core.roles import AccountKind, Role +from app.services.prediction import PredictionSummary, SitePrediction, SitePredictionSummary + +TARGET_AT = datetime(2026, 9, 16, 13, 0, tzinfo=UTC) +CREATED_AT = datetime(2026, 9, 16, 12, 0, tzinfo=UTC) + + +def principal(role: Role = Role.LECTEUR) -> Principal: + return Principal( + id=uuid4(), + email=f"{role.value}@enervision.fr", + role=role, + kind=AccountKind.HUMAIN, + must_change_password=False, + ) + + +class FauxService: + def __init__(self) -> None: + self.resume = PredictionSummary( + timestamp=datetime.now(UTC), + sites=[ + SitePredictionSummary( + site_id="SITE001", + site_name="Bureau Paris La Défense", + prediction=SitePrediction( + target_at=TARGET_AT, + target_metric="consumption_kwh", + period_minutes=60, + predicted_value=812.5, + status="available", + failure_reason=None, + model_reference="lightgbm-abc123", + created_at=CREATED_AT, + ), + ), + SitePredictionSummary(site_id="SITE002", site_name="Usine Lyon", prediction=None), + ], + ) + + async def summary(self) -> PredictionSummary: + return self.resume + + +@pytest.fixture +def servi(app: FastAPI) -> Iterator[Callable[[], FauxService]]: + def installe() -> FauxService: + service = FauxService() + app.dependency_overrides[get_prediction_service] = lambda: service + app.dependency_overrides[get_current_principal] = lambda: principal() + return service + + yield installe + app.dependency_overrides.pop(get_prediction_service, None) + app.dependency_overrides.pop(get_current_principal, None) + + +async def test_get_predictions_returns_the_service_result( + servi: Callable[[], FauxService], client: AsyncClient +) -> None: + servi() + + response = await client.get("/api/v1/predictions") + + assert response.status_code == 200 + corps = response.json() + premier, second = corps["sites"] + assert premier["site_id"] == "SITE001" + assert premier["prediction"]["predicted_value"] == 812.5 + assert premier["prediction"]["status"] == "available" + assert second["site_id"] == "SITE002" + assert second["prediction"] is None diff --git a/apps/backend/tests/repositories/test_prediction.py b/apps/backend/tests/repositories/test_prediction.py new file mode 100644 index 0000000..2be7ddd --- /dev/null +++ b/apps/backend/tests/repositories/test_prediction.py @@ -0,0 +1,92 @@ +from datetime import UTC, datetime + +import pytest +from sqlalchemy.ext.asyncio import AsyncSession + +from app.models.energy import Prediction +from app.repositories.prediction import PredictionRepository +from tests.repositories.test_site import creer as creer_site +from tests.repositories.test_site import identifiant as identifiant_site + +pytestmark = pytest.mark.integration + + +async def creer_prediction( + session: AsyncSession, *, site_id: str, **overrides: object +) -> Prediction: + prediction = Prediction( + site_id=site_id, + target_at=overrides.get("target_at", datetime(2026, 9, 16, tzinfo=UTC)), + target_metric=overrides.get("target_metric", "consumption_kwh"), + period_minutes=overrides.get("period_minutes", 60), + predicted_value=overrides.get("predicted_value", 42.0), + model_reference=overrides.get("model_reference", "lightgbm-test"), + status=overrides.get("status", "available"), + failure_reason=overrides.get("failure_reason"), + ) + session.add(prediction) + await session.flush() + return prediction + + +async def test_latest_by_site_keeps_only_the_most_recent_target(session: AsyncSession) -> None: + site = await creer_site(session) + depot = PredictionRepository(session) + ancienne = await creer_prediction( + session, site_id=site.site_id, target_at=datetime(2026, 9, 1, tzinfo=UTC) + ) + recente = await creer_prediction( + session, site_id=site.site_id, target_at=datetime(2026, 9, 15, tzinfo=UTC) + ) + + resultats = await depot.latest_by_site() + identifiants = [ + p.prediction_id + for p in resultats + if p.prediction_id in (ancienne.prediction_id, recente.prediction_id) + ] + await session.rollback() + + assert identifiants == [recente.prediction_id] + + +async def test_latest_by_site_returns_one_row_per_site(session: AsyncSession) -> None: + premier = await creer_site(session) + second = await creer_site(session) + depot = PredictionRepository(session) + voulue_premier = await creer_prediction(session, site_id=premier.site_id) + voulue_second = await creer_prediction(session, site_id=second.site_id) + + resultats = await depot.latest_by_site() + identifiants = {p.site_id for p in resultats if p.site_id in (premier.site_id, second.site_id)} + await session.rollback() + + assert identifiants == {voulue_premier.site_id, voulue_second.site_id} + + +async def test_latest_by_site_keeps_an_insufficient_data_prediction(session: AsyncSession) -> None: + site = await creer_site(session) + depot = PredictionRepository(session) + voulue = await creer_prediction( + session, + site_id=site.site_id, + status="insufficient_data", + predicted_value=None, + failure_reason="pas assez d'historique", + ) + + resultats = await depot.latest_by_site() + identifiants = [p.prediction_id for p in resultats if p.site_id == site.site_id] + await session.rollback() + + assert identifiants == [voulue.prediction_id] + + +async def test_latest_by_site_returns_an_empty_list_when_there_is_nothing( + session: AsyncSession, +) -> None: + depot = PredictionRepository(session) + + resultats = [p for p in await depot.latest_by_site() if p.site_id == identifiant_site()] + + assert resultats == [] diff --git a/apps/backend/tests/services/test_prediction.py b/apps/backend/tests/services/test_prediction.py new file mode 100644 index 0000000..a402b24 --- /dev/null +++ b/apps/backend/tests/services/test_prediction.py @@ -0,0 +1,121 @@ +from dataclasses import dataclass +from datetime import UTC, datetime + +from app.services.prediction import PredictionService + +TARGET_AT = datetime(2026, 9, 16, 13, 0, tzinfo=UTC) +CREATED_AT = datetime(2026, 9, 16, 12, 0, tzinfo=UTC) + + +@dataclass +class FauxSite: + site_id: str + site_name: str + + +@dataclass +class FauxPrediction: + site_id: str + target_at: datetime + target_metric: str + period_minutes: int | None + predicted_value: float | None + status: str + failure_reason: str | None + model_reference: str + created_at: datetime + + +class FauxDepotSites: + def __init__(self, sites: list[FauxSite]) -> None: + self._sites = sites + + async def list_all(self) -> list[FauxSite]: + return self._sites + + +class FauxDepotPredictions: + def __init__(self, predictions: list[FauxPrediction]) -> None: + self._predictions = predictions + + async def latest_by_site(self) -> list[FauxPrediction]: + return self._predictions + + +def prediction_disponible(site_id: str = "A") -> FauxPrediction: + return FauxPrediction( + site_id=site_id, + target_at=TARGET_AT, + target_metric="consumption_kwh", + period_minutes=60, + predicted_value=812.5, + status="available", + failure_reason=None, + model_reference="lightgbm-abc123", + created_at=CREATED_AT, + ) + + +async def test_summary_attaches_the_latest_prediction_to_its_site() -> None: + service = PredictionService( + sites=FauxDepotSites([FauxSite("A", "Site A")]), # type: ignore[arg-type] + predictions=FauxDepotPredictions([prediction_disponible("A")]), # type: ignore[arg-type] + ) + + resume = await service.summary() + + site = resume.sites[0] + assert site.site_id == "A" + assert site.prediction is not None + assert site.prediction.predicted_value == 812.5 + assert site.prediction.status == "available" + + +async def test_summary_leaves_prediction_none_for_a_site_never_scored() -> None: + service = PredictionService( + sites=FauxDepotSites([FauxSite("A", "Site A")]), # type: ignore[arg-type] + predictions=FauxDepotPredictions([]), # type: ignore[arg-type] + ) + + resume = await service.summary() + + assert resume.sites[0].prediction is None + + +async def test_summary_carries_an_insufficient_data_prediction_without_a_value() -> None: + insuffisante = FauxPrediction( + site_id="A", + target_at=TARGET_AT, + target_metric="consumption_kwh", + period_minutes=60, + predicted_value=None, + status="insufficient_data", + failure_reason="pas assez d'historique", + model_reference="lightgbm-abc123", + created_at=CREATED_AT, + ) + service = PredictionService( + sites=FauxDepotSites([FauxSite("A", "Site A")]), # type: ignore[arg-type] + predictions=FauxDepotPredictions([insuffisante]), # type: ignore[arg-type] + ) + + resume = await service.summary() + + site = resume.sites[0] + assert site.prediction is not None + assert site.prediction.status == "insufficient_data" + assert site.prediction.predicted_value is None + assert site.prediction.failure_reason == "pas assez d'historique" + + +async def test_summary_covers_every_site_even_with_a_single_prediction_in_the_repository() -> None: + service = PredictionService( + sites=FauxDepotSites([FauxSite("A", "Site A"), FauxSite("B", "Site B")]), # type: ignore[arg-type] + predictions=FauxDepotPredictions([prediction_disponible("A")]), # type: ignore[arg-type] + ) + + resume = await service.summary() + + par_site = {site.site_id: site for site in resume.sites} + assert par_site["A"].prediction is not None + assert par_site["B"].prediction is None diff --git a/docs/architecture/00-vue-ensemble.md b/docs/architecture/00-vue-ensemble.md index e2b68b1..26b890e 100644 --- a/docs/architecture/00-vue-ensemble.md +++ b/docs/architecture/00-vue-ensemble.md @@ -74,10 +74,10 @@ collecteur ne vient le lire. | Domaine | Technologie | Emplacement | Statut | Ce qui existe réellement | |---|---|---|---|---| -| 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) | +| 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`, `readings`, `sensors/status` et `predictions` 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 | +| ML | LightGBM, MLflow | `ml` | `En cours` | Pipeline d'entraînement et de scoring (`enervision_ml.train`/`.score`, features par lags/moyennes glissantes partagées entre les deux, baseline de persistance saisonnière, suivi MLflow local), exposé en lecture via `GET /predictions`. Voir [ADR 0005](../adr/0005-modele-prediction-lightgbm.md) et [ML-START.md](../../ML-START.md). Automatisation (Airflow) et surveillance de dérive (EC06, #44/#45) pas encore construites | | 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 | diff --git a/docs/architecture/20-backend.md b/docs/architecture/20-backend.md index 553ac4e..2e5cc3b 100644 --- a/docs/architecture/20-backend.md +++ b/docs/architecture/20-backend.md @@ -12,10 +12,10 @@ Les quatre couches existent désormais, portées par l'authentification. ```mermaid flowchart TB - ep["endpoints
health, auth, users, sites, alerts,
recommendations, stats, sensors"] + ep["endpoints
health, auth, users, sites, alerts,
recommendations, stats, readings, sensors, predictions"] sc["schemas
Pydantic"] - sv["services
AuthService, UserService,
SiteService, AlertService, RecommendationService,
StatsService, SensorService"] - rp["repositories
user, refresh_token,
login_attempt, audit_log,
site, alert, recommendation, reading"] + sv["services
AuthService, UserService,
SiteService, AlertService, RecommendationService,
StatsService, ReadingService, SensorService, PredictionService"] + rp["repositories
user, refresh_token,
login_attempt, audit_log,
site, alert, recommendation, reading, prediction"] md["models
10 tables"] db[("PostgreSQL")] @@ -148,6 +148,7 @@ Deux fichiers d'environnement, deux usages : `.env` à la racine alimente `docke | GET | `/api/v1/stats/summary` | Résume la consommation instantanée du parc. `lecteur` | 401, 403, 500 | | GET | `/api/v1/readings` | Historique des lectures, filtrable par `site_id`, fenêtre `start`/`end` (24h par défaut, 90 jours maximum) et paginé par `limit`/`offset`. `lecteur` | 400, 401, 403, 422, 500 | | GET | `/api/v1/sensors/status` | État de santé des capteurs par site, dérivé de la dernière lecture. `admin` | 401, 403, 500 | +| GET | `/api/v1/predictions` | Dernière prévision de consommation par site, calculée hors ligne par le pipeline de scoring (`ml/`). `lecteur` | 401, 403, 500 | | GET | `/metrics` | Format Prometheus, hors du schéma. Jeton requis si `APP_METRICS_TOKEN` est posé | | | GET | `/docs`, `/redoc`, `/openapi.json` | Hors du schéma. Fermés en `staging` et en `prod` | | @@ -160,7 +161,7 @@ Les codes de la dernière colonne sont ceux que le schéma **déclare**, et le f donc de modifier la liste dans ce fichier de test. `GET /sites` et `GET /sites/{site_id}` sont la première route métier, et le gabarit repris pour -`GET /alerts` puis pour les suivantes (`dataset`, `prediction`) : les quatre couches +`GET /alerts` puis pour les suivantes (`dataset`) : les quatre couches `endpoints → services → repositories → models` y sont toutes présentes, sur des tables déjà créées par la révision Alembic `e6d2026091501`. Elles n'exigent que le rôle `lecteur`, contrairement aux routes d'administration qui exigent `admin`. `SiteRepository` lit par `AsyncSession.scalar()` (une @@ -174,6 +175,31 @@ elle remonte à un site par sa seule `alert_id`, `alert` n'étant pas encore exp pas dans ce gabarit route-par-table. Le contrat détaillé pour le frontend est dans [31-contrat-authentification.md](31-contrat-authentification.md). +`GET /predictions` reprend ce même sous-gabarit « dernière valeur par site » (`SiteRepository` + +`PredictionRepository`, un `SitePredictionSummaryResponse` par site plutôt qu'une table brute). +Différence avec `stats`/`sensors` : `prediction` est une vraie table accumulée par un processus +externe (`enervision_ml.score`, cf. `ml/README.md`), pas une valeur recalculée à la volée depuis +`reading` à chaque appel. `PredictionRepository.latest_by_site()` isole donc un `DISTINCT ON +(site_id)` ordonné par `target_at DESC` (couvert par l'index `ix_prediction_site_target`), le même +mécanisme que `ReadingRepository.latest_by_site()`. Un site jamais scoré rend `prediction: null` +plutôt qu'un statut inventé : le domaine `available`/`insufficient_data`/`error` de la contrainte +`ck_prediction_status` n'a pas de valeur pour « pas encore de ligne ». L'API ne lance jamais +LightGBM elle-même ; elle lit ce que le pipeline de scoring a déjà écrit, cf. +[ML-START.md](../../ML-START.md) section 3. + +`GET /readings` reprend le même gabarit mais s'en écarte sur un point : `reading` est l'hypertable, +donc la seule table métier pouvant porter des années d'historique, ce que `docs/architecture/ +owasp-traceabilite.md` documentait comme un risque ouvert (API4, aucune pagination plafonnée ni +fenêtre temporelle maximale). `ReadingService` porte donc une couche de validation absente des +autres routes de lecture : `start`/`end` sont optionnels (24 dernières heures par défaut si les +deux sont omis, l'un défaut par rapport à l'autre sinon), l'écart entre les deux est plafonné à 90 +jours (`FENETRE_MAXIMALE`), et `limit`/`offset` (défaut 500, plafond 2000) empêchent qu'une fenêtre +large mais peu dense reste malgré tout coûteuse. Un dépassement de plafond répond `400` (règle +métier, portée par le service) plutôt que `422` (réservé à la validation structurelle de FastAPI, +par exemple `limit` hors bornes). Un datetime sans fuseau dans `start`/`end` est traité comme de +l'UTC plutôt que rejeté : le comparer tel quel à `reading.timestamp` (`timestamptz`) échouerait +côté pilote, en `500` plutôt qu'un refus propre. + `GET /readings` reprend le même gabarit mais s'en écarte sur un point : `reading` est l'hypertable, donc la seule table métier pouvant porter des années d'historique, ce que `docs/architecture/ owasp-traceabilite.md` documentait comme un risque ouvert (API4, aucune pagination plafonnée ni @@ -262,7 +288,7 @@ Les modèles de `app/schemas/errors.py` décrivent ce que les gestionnaires renv ### Ajouter une route métier Checklist pour toute nouvelle route sur le gabarit `sites`/`alerts`/`recommendations`/`stats`/ -`readings`/`sensors` (`dataset`, `prediction`) : +`readings`/`sensors`/`predictions` (`dataset`) : 1. Composer ses `responses=` depuis `app/api/openapi.py` : `REPONSES_LECTEUR`/`REPONSES_ADMIN` au niveau de l'`include_router()` du routeur, `REPONSE_VALIDATION` et les codes locaux diff --git a/ml/README.md b/ml/README.md index 4d69363..c7814fe 100644 --- a/ml/README.md +++ b/ml/README.md @@ -57,6 +57,46 @@ validation. La coupure est **chronologique**, jamais un tirage aleatoire de lign aleatoire laisserait des lignes de validation "voir" des lignes d'entrainement via leurs lags/moyennes glissantes, une fuite qui masquerait un surapprentissage. +## Scoring + +```bash +uv run python -m enervision_ml.score --csv data/all_sites_combined.csv +# ou, une fois la base peuplee et ML_DATABASE_URL positionnee : +uv run python -m enervision_ml.score +``` + +Calcule, pour chaque site (ou un seul avec `--site-id`), la consommation prevue de l'heure suivant +sa derniere lecture connue, et ecrit une ligne dans `prediction`. Etapes, cf. `ML-START.md` +section 2 : + +1. Lit une fenetre recente de `reading`+`site` (21 jours par defaut, une marge au-dessus des 168h + necessaires au lag hebdomadaire) plutot que tout l'historique -- le meme piege que celui deja + corrige sur `GET /readings` (fenetre non plafonnee sur une hypertable). +2. Ajoute une ligne "future" par site (l'heure suivante) et calcule ses features avec + `enervision_ml.features.build_features`, **exactement** la meme fonction qu'a l'entrainement. +3. Si le lag de 168h est absent (moins d'une semaine d'historique pour ce site) : ecrit + `status="insufficient_data"` directement, sans jamais appeler LightGBM. +4. Sinon : appelle `booster.predict(...)` et ecrit `status="available"` avec la valeur predite. + +`--model` pointe vers le fichier entraine (`models/lightgbm-consumption.txt` par defaut). +`model_reference` en base est le hache SHA-256 (tronque) du fichier modele, pas son nom de +fichier : `train.py` reecrit toujours le meme chemin a chaque entrainement, donc le nom seul ne +distinguerait pas deux versions du modele. + +En mode `--csv`, rien n'est ecrit en base : c'est un instantane historique fige (l'heure "future" +calculee a partir de la fin du CSV n'existe dans aucune base reelle), utile pour valider le +pipeline sans base joignable. + +**Limite assumee** : la feature `is_working_hours` de la ligne future est recopiee depuis la +derniere lecture reelle, pas recalculee -- il n'existe aucune regle horaire ouvrable dans ce +depot (elle vit dans le generateur du jeu de donnees d'origine). L'approximation n'est fausse +qu'aux heures de bascule ouverture/fermeture, sur une seule feature parmi une dizaine, pour une +prevision a un seul pas. + +`prediction` n'a pas de contrainte d'unicite sur `(site_id, target_at)` : chaque run de scoring +insere une nouvelle ligne plutot que d'ecraser la precedente, pour garder une trace de chaque +prevision (utile plus tard pour comparer prevision et realise, surveillance de derive #44/#45). + ## Commandes ```bash @@ -81,8 +121,14 @@ 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`. +du modele, a l'entrainement comme au scoring (`enervision_ml.score`). 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`. + +## Et cote API ? + +`GET /api/v1/predictions` (backend, `apps/backend`) lit ce que `enervision_ml.score` a ecrit dans +`prediction` -- la derniere prevision par site, jamais un recalcul a la volee. FastAPI ne fait +jamais tourner LightGBM lui-meme, cf. `ML-START.md` section 3. diff --git a/ml/enervision_ml/data.py b/ml/enervision_ml/data.py index e7b5853..6e08b54 100644 --- a/ml/enervision_ml/data.py +++ b/ml/enervision_ml/data.py @@ -16,6 +16,7 @@ Deux chemins, qui doivent produire le meme schema de sortie (colonnes `site_id`, colonne est renvoyee a `NaN`, que LightGBM gere nativement comme valeur manquante. """ +from datetime import datetime from pathlib import Path import pandas as pd @@ -34,6 +35,14 @@ OUTPUT_COLUMNS = [ "capacity_kw", ] +NUMERIC_COLUMNS = [ + "consumption_kwh", + "temperature_celsius", + "humidity_percent", + "solar_irradiance_wm2", + "capacity_kw", +] + _READING_QUERY = text( """ SELECT @@ -53,10 +62,43 @@ _READING_QUERY = text( ) +_RECENT_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 + WHERE r.timestamp >= :since + ORDER BY r.site_id, r.timestamp + """ +) + + def load_from_database(connection: Connectable) -> pd.DataFrame: - """Lit l'historique complet `reading` + `site` depuis PostgreSQL.""" + """Lit l'historique complet `reading` + `site` depuis PostgreSQL. Entrainement seulement : + le scoring n'a besoin que d'une fenetre recente, cf. `load_recent_from_database`. + """ frame = pd.read_sql(_READING_QUERY, connection) - return frame[OUTPUT_COLUMNS] + return _typer(frame[OUTPUT_COLUMNS]) + + +def load_recent_from_database(connection: Connectable, *, since: datetime) -> pd.DataFrame: + """Lit `reading` + `site` depuis `since` seulement, pour le scoring. + + Piege evite : un `SELECT` sans borne sur l'hypertable complete juste pour scorer le prochain + pas horaire serait la meme erreur que celle corrigee sur `GET /readings` (fenetre non + plafonnee sur une table pouvant porter des annees d'historique). + """ + frame = pd.read_sql(_RECENT_READING_QUERY, connection, params={"since": since}) + return _typer(frame[OUTPUT_COLUMNS]) def load_from_csv(csv_path: Path) -> pd.DataFrame: @@ -65,4 +107,20 @@ def load_from_csv(csv_path: Path) -> pd.DataFrame: frame["capacity_kw"] = float("nan") frame["is_working_hours"] = frame["is_working_hours"].astype(bool) - return frame[OUTPUT_COLUMNS] + return _typer(frame[OUTPUT_COLUMNS]) + + +def _typer(frame: pd.DataFrame) -> pd.DataFrame: + """Force le typage numerique attendu par LightGBM. + + Piege reel, pas theorique : `site.capacity_kw` n'est peuple par aucun pipeline d'ingestion + aujourd'hui (`historical_import.py` ne pose que `site_type`/`site_name`). Une colonne + entierement `NULL` revient de `pd.read_sql` en dtype `object` plutot que `float64`, ce que + LightGBM refuse ("pandas dtypes must be int, float or bool"). `pd.to_numeric` corrige aussi + n'importe quelle autre colonne mesuree entierement absente sur une fenetre de scoring, pas + seulement `capacity_kw`. + """ + typee = frame.copy() + for colonne in NUMERIC_COLUMNS: + typee[colonne] = pd.to_numeric(typee[colonne], errors="coerce") + return typee diff --git a/ml/enervision_ml/score.py b/ml/enervision_ml/score.py new file mode 100644 index 0000000..adf9634 --- /dev/null +++ b/ml/enervision_ml/score.py @@ -0,0 +1,281 @@ +"""Scoring du modele LightGBM : calcule et enregistre la consommation prevue du prochain pas +horaire, par site. + +CLI autonome, sur le meme gabarit que `enervision_ml.train` et +`apps/backend/app/etl/historical_import.py`. Cf. `docs/ML-START.md`, section 2. + + uv run python -m enervision_ml.score --csv ../ml/data/all_sites_combined.csv + uv run python -m enervision_ml.score # lit ML_DATABASE_URL, ecrit dans `prediction` + +Reutilise `enervision_ml.features.build_features` tel quel (jamais reecrit) : c'est la garantie +contre le train/serve skew documentee dans ce module. +""" + +import argparse +import hashlib +from dataclasses import dataclass +from datetime import UTC, datetime, timedelta +from pathlib import Path +from typing import Any, cast + +import lightgbm as lgb +import pandas as pd +from sqlalchemy import create_engine, text +from sqlalchemy.engine import Connection + +from enervision_ml import config +from enervision_ml.data import load_from_csv, load_recent_from_database +from enervision_ml.features import TARGET_COLUMN, WEATHER_COLUMNS, build_features, feature_columns + +# Marge au-dessus des 168h necessaires au lag hebdomadaire, pour absorber les trous de mesure. +LOOKBACK = timedelta(days=21) + +TARGET_METRIC = "consumption_kwh" +PERIOD_MINUTES = 60 +LAG_168H_COLUMN = f"{TARGET_COLUMN}_lag_168h" +INSUFFICIENT_DATA_REASON = ( + "Historique insuffisant : moins de 168h de consumption_kwh disponibles pour ce site." +) + + +@dataclass(frozen=True, slots=True) +class ScoredSite: + site_id: str + target_at: datetime + status: str + predicted_value: float | None + failure_reason: str | None + + +def model_reference(model_path: Path) -> str: + """Identifiant stable du modele utilise, insensible au fait que `train.py` reecrive + toujours le meme nom de fichier a chaque entrainement (pas de versioning par nom, cf. + `ml/README.md`).""" + empreinte = hashlib.sha256(model_path.read_bytes()).hexdigest() + return f"lightgbm-{empreinte[:12]}" + + +def build_scoring_frame(recent: pd.DataFrame, *, site_id: str | None = None) -> pd.DataFrame: + """Ajoute une ligne future (l'heure suivant la derniere lecture connue) par site, et calcule + ses features par `build_features` -- exactement comme a l'entrainement, seule la cible de + cette ligne est inconnue. + + Piege assume : `is_working_hours` de la ligne future est copie de la derniere lecture reelle, + pas recalcule. Il n'existe aucune regle horaire ouvrable dans ce depot (elle vit dans le + generateur du jeu de donnees d'origine, hors de ce code) ; l'approximation n'est fausse + qu'aux heures de bascule (ouverture/fermeture), sur une seule feature parmi une dizaine, pour + une prevision a un pas seulement. + """ + travail = recent if site_id is None else recent[recent["site_id"] == site_id] + if travail.empty: + return build_features(travail) + + dernieres = ( + travail.sort_values("timestamp").groupby("site_id", as_index=False, sort=False).tail(1) + ).copy() + dernieres["timestamp"] = dernieres["timestamp"] + pd.Timedelta(hours=1) + dernieres[TARGET_COLUMN] = float("nan") + # Meteo future inconnue (cf. piege documente dans `enervision_ml.features.build_features`) : + # laisser `NaN` ici n'a aucun effet sur les features utilisees, qui ne prennent la meteo que + # decalee. + for colonne in WEATHER_COLUMNS: + dernieres[colonne] = float("nan") + + etendu = pd.concat([travail, dernieres], ignore_index=True) + features = build_features(etendu) + return features.groupby("site_id", as_index=False, sort=False).tail(1).reset_index(drop=True) + + +def score(booster: lgb.Booster, scoring_frame: pd.DataFrame) -> list[ScoredSite]: + resultats: list[ScoredSite] = [] + + insuffisants = scoring_frame[scoring_frame[LAG_168H_COLUMN].isna()] + for enregistrement in _records(insuffisants): + resultats.append( + ScoredSite( + site_id=enregistrement["site_id"], + target_at=enregistrement["timestamp"].to_pydatetime(), + status="insufficient_data", + predicted_value=None, + failure_reason=INSUFFICIENT_DATA_REASON, + ) + ) + + suffisants = scoring_frame[scoring_frame[LAG_168H_COLUMN].notna()] + if not suffisants.empty: + typee = suffisants.copy() + typee["site_type"] = typee["site_type"].astype("category") + predictions = booster.predict(typee[feature_columns()]) + for enregistrement, valeur in zip(_records(suffisants), predictions, strict=True): + resultats.append( + ScoredSite( + site_id=enregistrement["site_id"], + target_at=enregistrement["timestamp"].to_pydatetime(), + status="available", + predicted_value=float(valeur), + failure_reason=None, + ) + ) + + return resultats + + +def _records(frame: pd.DataFrame) -> list[dict[str, Any]]: + return cast(list[dict[str, Any]], frame.to_dict(orient="records")) + + +_INSERT_PREDICTION = text( + """ + INSERT INTO prediction ( + site_id, target_at, target_metric, period_minutes, + predicted_value, model_reference, status, failure_reason + ) VALUES ( + :site_id, :target_at, :target_metric, :period_minutes, + :predicted_value, :model_reference, :status, :failure_reason + ) + """ +) + + +def write_predictions( + connection: Connection, resultats: list[ScoredSite], *, reference: str +) -> None: + """Ecrit une ligne par site score. Insertion seule, jamais de mise a jour : `prediction` + n'a pas de contrainte d'unicite sur `(site_id, target_at)`, chaque run garde sa propre trace + plutot que d'ecraser la precedente -- utile plus tard pour comparer prevision et realise + (surveillance de derive, #44/#45).""" + if not resultats: + return + + lignes = [ + { + "site_id": r.site_id, + "target_at": r.target_at, + "target_metric": TARGET_METRIC, + "period_minutes": PERIOD_MINUTES, + "predicted_value": r.predicted_value, + "model_reference": reference, + "status": r.status, + "failure_reason": r.failure_reason, + } + for r in resultats + ] + connection.execute(_INSERT_PREDICTION, lignes) + + +def _load_recent(*, csv_path: Path | None, now: datetime | None) -> tuple[pd.DataFrame, datetime]: + if csv_path is not None: + brute = load_from_csv(csv_path) + instant = now or ( + brute["timestamp"].max().to_pydatetime() if not brute.empty else datetime.now(UTC) + ) + return brute[brute["timestamp"] >= instant - LOOKBACK], instant + + instant = now or datetime.now(UTC) + engine = create_engine(config.database_url()) + try: + return load_recent_from_database(engine, since=instant - LOOKBACK), instant + finally: + engine.dispose() + + +def run_scoring( + *, + model_path: Path, + csv_path: Path | None = None, + site_id: str | None = None, + now: datetime | None = None, +) -> list[ScoredSite]: + """Score le prochain pas horaire par site et l'ecrit dans `prediction`. + + En mode `--csv`, rien n'est ecrit : c'est un instantane historique fige (l'heure "future" + calculee n'existe dans aucune base reelle), utile pour valider le pipeline sans base + joignable, cf. `ml/README.md`. + """ + recent, _instant = _load_recent(csv_path=csv_path, now=now) + if site_id is not None: + recent = recent[recent["site_id"] == site_id] + + scoring_frame = build_scoring_frame(recent, site_id=site_id) + if scoring_frame.empty: + return [] + + booster = lgb.Booster(model_file=str(model_path)) + resultats = score(booster, scoring_frame) + + if csv_path is None: + reference = model_reference(model_path) + engine = create_engine(config.database_url()) + try: + with engine.begin() as connection: + write_predictions(connection, resultats, reference=reference) + finally: + engine.dispose() + + return resultats + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Scoring du modele LightGBM EnerVision") + + parser.add_argument( + "--model", + type=Path, + default=Path("models/lightgbm-consumption.txt"), + help="Chemin du modele entraine. Defaut : models/lightgbm-consumption.txt.", + ) + parser.add_argument( + "--csv", + type=Path, + default=None, + help=( + "Instantane historique de demarrage/demo, rien n'est ecrit en base. Omis, lit " + "ML_DATABASE_URL, se connecte a PostgreSQL et ecrit dans `prediction`." + ), + ) + parser.add_argument( + "--site-id", + default=None, + help="Ne score que ce site. Omis, tous les sites presents dans la fenetre recente.", + ) + parser.add_argument( + "--now", + type=_parse_instant, + default=None, + help=( + "Instant de reference (ISO 8601), pour tester ou demontrer le scoring cote base sur " + "des donnees anciennes (ex. le jeu de donnees historique, qui s'arrete fin 2024). " + "Omis, horloge systeme reelle." + ), + ) + + return parser.parse_args() + + +def _parse_instant(valeur: str) -> datetime: + instant = datetime.fromisoformat(valeur) + return instant if instant.tzinfo is not None else instant.replace(tzinfo=UTC) + + +def main() -> None: + args = parse_args() + resultats = run_scoring( + model_path=args.model, csv_path=args.csv, site_id=args.site_id, now=args.now + ) + + if not resultats: + print("Aucun site a scorer (aucune lecture recente dans la fenetre).") + return + + for r in resultats: + if r.status == "available": + print(f"{r.site_id} @ {r.target_at} : {r.predicted_value:.2f} kWh") + else: + print(f"{r.site_id} @ {r.target_at} : {r.status} ({r.failure_reason})") + + if args.csv is not None: + print("\nMode --csv : instantane historique, rien ecrit en base.") + + +if __name__ == "__main__": + main() diff --git a/ml/pyproject.toml b/ml/pyproject.toml index b3c368b..1589c92 100644 --- a/ml/pyproject.toml +++ b/ml/pyproject.toml @@ -47,9 +47,9 @@ select = [ "S", "PT", ] -# N806 : `X`/`y` (donnees/cible) est la convention scikit-learn/LightGBM, pas une variable mal -# nommee. -ignore = ["B008", "N806"] +# N806/N803 : `X`/`y` (donnees/cible) est la convention scikit-learn/LightGBM, pas une variable +# ou un argument mal nomme. +ignore = ["B008", "N806", "N803"] [tool.ruff.lint.per-file-ignores] "tests/**/*.py" = ["S101"] diff --git a/ml/tests/test_data.py b/ml/tests/test_data.py new file mode 100644 index 0000000..0a1f022 --- /dev/null +++ b/ml/tests/test_data.py @@ -0,0 +1,34 @@ +import pandas as pd + +from enervision_ml.data import NUMERIC_COLUMNS, OUTPUT_COLUMNS, _typer + + +def make_frame_with_object_dtype_capacity() -> pd.DataFrame: + # Reproduit ce que `pd.read_sql` renvoie pour une colonne entierement `NULL` en base : + # dtype `object` rempli de `None`, pas `float64` rempli de `NaN`. + frame = pd.DataFrame( + {colonne: [1.0, 2.0] for colonne in OUTPUT_COLUMNS if colonne not in NUMERIC_COLUMNS} + ) + for colonne in NUMERIC_COLUMNS: + frame[colonne] = pd.Series([None, None], dtype="object") + return frame + + +def test_typer_coerces_an_all_null_object_column_to_float() -> None: + frame = make_frame_with_object_dtype_capacity() + + typee = _typer(frame) + + for colonne in NUMERIC_COLUMNS: + assert typee[colonne].dtype == "float64" + assert typee[colonne].isna().all() + + +def test_typer_preserves_real_numeric_values() -> None: + frame = make_frame_with_object_dtype_capacity() + frame["capacity_kw"] = pd.Series([100.0, None], dtype="object") + + typee = _typer(frame) + + assert typee["capacity_kw"].tolist()[0] == 100.0 + assert pd.isna(typee["capacity_kw"].tolist()[1]) diff --git a/ml/tests/test_score.py b/ml/tests/test_score.py new file mode 100644 index 0000000..28b0af5 --- /dev/null +++ b/ml/tests/test_score.py @@ -0,0 +1,230 @@ +from datetime import UTC, datetime, timedelta +from pathlib import Path +from typing import Any + +import pandas as pd +import pytest + +from enervision_ml.features import TARGET_COLUMN +from enervision_ml.score import ( + LAG_168H_COLUMN, + ScoredSite, + build_scoring_frame, + model_reference, + run_scoring, + score, + write_predictions, +) + + +def make_recent( + 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, + "humidity_percent": 50.0, + "solar_irradiance_wm2": 0.0, + "is_working_hours": True, + "site_type": "office", + "capacity_kw": 100.0, + } + ) + + +class FakeBooster: + def __init__(self, valeur: float = 42.0) -> None: + self.valeur = valeur + self.appels: list[int] = [] + + def predict(self, X: Any) -> list[float]: + self.appels.append(len(X)) + return [self.valeur] * len(X) + + +class FakeConnection: + def __init__(self) -> None: + self.appels: list[tuple[Any, Any]] = [] + + def execute(self, statement: Any, parameters: Any = None) -> None: + self.appels.append((statement, parameters)) + + +def test_build_scoring_frame_adds_one_row_per_site_one_hour_after_the_last_reading() -> None: + depart = datetime(2026, 1, 1, tzinfo=UTC) + recent = pd.concat( + [ + make_recent("site-a", heures=200, depart=depart), + make_recent("site-b", heures=200, depart=depart), + ], + ignore_index=True, + ) + + scoring_frame = build_scoring_frame(recent) + + assert set(scoring_frame["site_id"]) == {"site-a", "site-b"} + derniere_lecture = depart + timedelta(hours=199) + assert (scoring_frame["timestamp"] == derniere_lecture + timedelta(hours=1)).all() + + +def test_build_scoring_frame_computes_lags_from_real_history() -> None: + depart = datetime(2026, 1, 1, tzinfo=UTC) + recent = make_recent("site-a", heures=200, depart=depart) + + scoring_frame = build_scoring_frame(recent) + + ligne = scoring_frame.iloc[0] + # La cible future n'existe pas : le lag d'1h doit valoir la toute derniere valeur reelle. + assert ligne[f"{TARGET_COLUMN}_lag_1h"] == recent[TARGET_COLUMN].iloc[-1] + + +def test_build_scoring_frame_flags_insufficient_history_under_168_hours() -> None: + depart = datetime(2026, 1, 1, tzinfo=UTC) + recent = make_recent("site-a", heures=100, depart=depart) + + scoring_frame = build_scoring_frame(recent) + + assert pd.isna(scoring_frame.iloc[0][LAG_168H_COLUMN]) + + +def test_build_scoring_frame_accepts_a_full_week_of_history() -> None: + depart = datetime(2026, 1, 1, tzinfo=UTC) + recent = make_recent("site-a", heures=169, depart=depart) + + scoring_frame = build_scoring_frame(recent) + + assert not pd.isna(scoring_frame.iloc[0][LAG_168H_COLUMN]) + + +def test_build_scoring_frame_filters_to_a_single_site() -> None: + depart = datetime(2026, 1, 1, tzinfo=UTC) + recent = pd.concat( + [ + make_recent("site-a", heures=200, depart=depart), + make_recent("site-b", heures=200, depart=depart), + ], + ignore_index=True, + ) + + scoring_frame = build_scoring_frame(recent, site_id="site-a") + + assert scoring_frame["site_id"].tolist() == ["site-a"] + + +def test_build_scoring_frame_returns_empty_when_there_is_no_recent_reading() -> None: + recent = make_recent("site-a", heures=0, depart=datetime(2026, 1, 1, tzinfo=UTC)) + + scoring_frame = build_scoring_frame(recent) + + assert scoring_frame.empty + + +def test_score_marks_insufficient_history_without_calling_the_model() -> None: + depart = datetime(2026, 1, 1, tzinfo=UTC) + scoring_frame = build_scoring_frame(make_recent("site-a", heures=100, depart=depart)) + booster = FakeBooster() + + resultats = score(booster, scoring_frame) # type: ignore[arg-type] + + assert resultats == [ + ScoredSite( + site_id="site-a", + target_at=resultats[0].target_at, + status="insufficient_data", + predicted_value=None, + failure_reason=resultats[0].failure_reason, + ) + ] + assert booster.appels == [] + + +def test_score_predicts_when_history_is_sufficient() -> None: + depart = datetime(2026, 1, 1, tzinfo=UTC) + scoring_frame = build_scoring_frame(make_recent("site-a", heures=200, depart=depart)) + booster = FakeBooster(valeur=99.5) + + resultats = score(booster, scoring_frame) # type: ignore[arg-type] + + assert len(resultats) == 1 + assert resultats[0].status == "available" + assert resultats[0].predicted_value == 99.5 + assert resultats[0].failure_reason is None + assert booster.appels == [1] + + +def test_write_predictions_does_nothing_when_there_is_nothing_to_write() -> None: + connection = FakeConnection() + + write_predictions(connection, [], reference="lightgbm-test") # type: ignore[arg-type] + + assert connection.appels == [] + + +def test_write_predictions_sends_one_row_per_result() -> None: + connection = FakeConnection() + resultats = [ + ScoredSite("site-a", datetime(2026, 1, 1, tzinfo=UTC), "available", 42.0, None), + ScoredSite( + "site-b", + datetime(2026, 1, 1, tzinfo=UTC), + "insufficient_data", + None, + "pas assez d'historique", + ), + ] + + write_predictions(connection, resultats, reference="lightgbm-test") # type: ignore[arg-type] + + assert len(connection.appels) == 1 + _, lignes = connection.appels[0] + assert len(lignes) == 2 + assert lignes[0]["model_reference"] == "lightgbm-test" + assert lignes[0]["target_metric"] == "consumption_kwh" + assert lignes[0]["period_minutes"] == 60 + + +def test_model_reference_is_stable_for_the_same_file_content(tmp_path: Path) -> None: + model_path = tmp_path / "model.txt" + model_path.write_bytes(b"contenu-du-modele") + + assert model_reference(model_path) == model_reference(model_path) + + +def test_model_reference_changes_with_the_file_content(tmp_path: Path) -> None: + premier = tmp_path / "model-a.txt" + premier.write_bytes(b"version-1") + second = tmp_path / "model-b.txt" + second.write_bytes(b"version-2") + + assert model_reference(premier) != model_reference(second) + + +def test_run_scoring_in_csv_mode_scores_without_touching_a_database(tmp_path: Path) -> None: + depart = datetime(2026, 1, 1, tzinfo=UTC) + frame = pd.concat( + [ + make_recent("site-a", heures=400, depart=depart), + make_recent("site-b", heures=400, depart=depart), + ], + ignore_index=True, + ) + csv_path = tmp_path / "recent.csv" + frame.to_csv(csv_path, index=False) + + model_path = tmp_path / "model.txt" + model_path.write_bytes(b"peu importe le contenu pour ce test") + + with pytest.MonkeyPatch.context() as monkeypatch: + monkeypatch.setattr( + "enervision_ml.score.lgb.Booster", lambda model_file: FakeBooster(valeur=7.0) + ) + + resultats = run_scoring(model_path=model_path, csv_path=csv_path) + + assert {r.site_id for r in resultats} == {"site-a", "site-b"} + assert all(r.status == "available" for r in resultats) + assert all(r.predicted_value == 7.0 for r in resultats)