feat(ml,backend): implemente le service de scoring et GET /predictions (#37)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
@@ -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]
|
||||
@@ -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
|
||||
)
|
||||
@@ -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`."
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
@@ -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"),
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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 == []
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user