feat(ml,backend): implemente le service de scoring et GET /predictions (#37)
This commit is contained in:
@@ -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