test(ml,backend): couvre ML vers DB, puis la chaine complete jusqu'a l'API
Le pipeline ML n'avait aucun test touchant PostgreSQL : `ml/README.md` le disait, faute de base joignable en CI. Le marqueur `integration` de `ml/pyproject.toml` etait declare et porte par zero test. - `ml/tests/conftest.py` : deux fixtures d'acces a la base, jamais interchangeables. `connexion_ml` annule sa transaction, `parc` valide ses ecritures parce que `run_scoring` ouvre sa propre connexion et ne verrait rien d'autre. Garde sur le nom de base, marque uuid sur chaque site, nettoyage dans l'ordre des cles etrangeres. - `test_data_integration.py` : les neuf colonnes du contrat confrontees au schema Alembic reel, la borne `since`, l'ordre de tri dont dependent des lags positionnels, et le typage des colonnes entierement nulles. - `test_score_integration.py` : les contraintes de `prediction` vues depuis le code qui ecrit, l'empilement volontaire de deux runs, et `run_scoring` de bout en bout sur un booster reel. - `apps/backend/tests/test_chaine_ml_api.py` : lance les vrais binaires `enervision_ml.train` et `.score` en sous-processus, comme les DAGs, puis relit par `GET /api/v1/predictions`. Marqueur `chaine` distinct : le job `integration` du backend n'a pas l'environnement de ml/. - `ml.yml` : job `integration`, seul du depot a reunir les deux environnements uv et une base. Ses `paths` incluent les migrations du backend, sans quoi le schema deriverait du SQL du pipeline sans que rien ne casse. - Makefile : `migrate-test`, qui manquait (`enervision_test` n'a jamais recu de table), `ml-test-integration` et `test-chaine`.
This commit is contained in:
@@ -0,0 +1,332 @@
|
||||
"""Piege : deux fixtures d'acces a la base, jamais interchangeables - `connexion_ml` et `parc`.
|
||||
|
||||
`connexion_ml` ouvre une transaction annulee a la fin du test : rien ne subsiste, et rien n'est
|
||||
visible hors de cette connexion. Elle sert aux fonctions qui recoivent leur connexion en
|
||||
argument (`load_from_database`, `load_recent_from_database`, `write_predictions`).
|
||||
|
||||
`run_scoring` fabrique en revanche son propre engine depuis `ML_DATABASE_URL` : il ne verrait
|
||||
pas des lignes semees dans une transaction non validee, et ses propres ecritures survivraient a
|
||||
l'annulation. Les tests qui l'appellent passent donc par `parc`, qui valide ce qu'il ecrit et
|
||||
nettoie lui-meme, dans l'ordre impose par les cles etrangeres `RESTRICT`.
|
||||
"""
|
||||
|
||||
import math
|
||||
import os
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
import lightgbm as lgb
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from sqlalchemy import Connection, Engine, Row, bindparam, create_engine, text
|
||||
from sqlalchemy.engine import URL, make_url
|
||||
|
||||
from enervision_ml.features import TARGET_COLUMN, build_features, feature_columns
|
||||
|
||||
BASE_ATTENDUE = "enervision_test"
|
||||
|
||||
# Piege : `load_from_database` lit toute la table, et `enervision_test` est partagee entre un run
|
||||
# local et la CI. Les tests ancrent donc leurs lectures au-dela de tout jeu de donnees reel
|
||||
# (l'historique s'arrete au 31/12/2024) pour que leur borne `since` ne ramene qu'eux.
|
||||
ANCRAGE = datetime(2035, 1, 1, tzinfo=UTC)
|
||||
|
||||
SITE_TYPE = "office"
|
||||
CAPACITY_KW = 100.0
|
||||
|
||||
_INSERT_SITE = text(
|
||||
"""
|
||||
INSERT INTO site (site_id, site_name, site_type, capacity_kw)
|
||||
VALUES (:site_id, :site_name, :site_type, :capacity_kw)
|
||||
"""
|
||||
)
|
||||
|
||||
# `source = 'api_history'` impose `dataset_id IS NULL` (ck_reading_dataset_source), ce qui evite
|
||||
# de creer une ligne `dataset`. `raw_data` est NOT NULL, d'ou le litteral jsonb.
|
||||
_INSERT_READING = text(
|
||||
"""
|
||||
INSERT INTO reading (
|
||||
site_id, timestamp, source, consumption_kwh, temperature_celsius,
|
||||
humidity_percent, solar_irradiance_wm2, is_working_hours, raw_data
|
||||
) VALUES (
|
||||
:site_id, :timestamp, :source, :consumption_kwh, :temperature_celsius,
|
||||
:humidity_percent, :solar_irradiance_wm2, :is_working_hours, '{}'::jsonb
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
_SELECT_PREDICTIONS = text(
|
||||
"""
|
||||
SELECT target_at, predicted_value, status, failure_reason, model_reference
|
||||
FROM prediction
|
||||
WHERE site_id = :site_id
|
||||
ORDER BY prediction_id
|
||||
"""
|
||||
)
|
||||
|
||||
_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, 'consumption_kwh', 60,
|
||||
:predicted_value, :model_reference, :status, :failure_reason
|
||||
)
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
# Ordre impose par les cles etrangeres `RESTRICT` : une lecture avant son site, une prediction
|
||||
# avant sa lecture.
|
||||
_SUPPRESSIONS = tuple(
|
||||
text(requete).bindparams(bindparam("sites", expanding=True))
|
||||
for requete in (
|
||||
"DELETE FROM prediction WHERE site_id IN :sites",
|
||||
"DELETE FROM reading WHERE site_id IN :sites",
|
||||
"DELETE FROM site WHERE site_id IN :sites",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def insere_site(
|
||||
connexion: Connection,
|
||||
*,
|
||||
site_type: str = SITE_TYPE,
|
||||
capacity_kw: float | None = CAPACITY_KW,
|
||||
) -> str:
|
||||
site_id = f"TEST-{uuid4().hex[:12]}"
|
||||
connexion.execute(
|
||||
_INSERT_SITE,
|
||||
{
|
||||
"site_id": site_id,
|
||||
"site_name": "Site de test",
|
||||
"site_type": site_type,
|
||||
"capacity_kw": capacity_kw,
|
||||
},
|
||||
)
|
||||
return site_id
|
||||
|
||||
|
||||
def insere_lectures(
|
||||
connexion: Connection,
|
||||
site_id: str,
|
||||
*,
|
||||
heures: int,
|
||||
fin: datetime,
|
||||
valeur: float = 50.0,
|
||||
source: str = "api_history",
|
||||
is_working_hours: bool | None = True,
|
||||
) -> list[datetime]:
|
||||
"""Grille horaire contigue finissant a `fin`, incluse.
|
||||
|
||||
Contigue parce que les lags de `build_features` sont des `shift()` positionnels : un trou
|
||||
dans la grille decalerait le lag de 168 h sans qu'aucune erreur ne se declenche.
|
||||
"""
|
||||
instants = [fin - timedelta(hours=decalage) for decalage in reversed(range(heures))]
|
||||
connexion.execute(
|
||||
_INSERT_READING,
|
||||
[
|
||||
{
|
||||
"site_id": site_id,
|
||||
"timestamp": instant,
|
||||
"source": source,
|
||||
"consumption_kwh": valeur + math.sin(rang / 12.0) * 10.0,
|
||||
"temperature_celsius": 15.0,
|
||||
"humidity_percent": 50.0,
|
||||
"solar_irradiance_wm2": 0.0,
|
||||
"is_working_hours": is_working_hours,
|
||||
}
|
||||
for rang, instant in enumerate(instants)
|
||||
],
|
||||
)
|
||||
return instants
|
||||
|
||||
|
||||
def insere_lecture(
|
||||
connexion: Connection,
|
||||
site_id: str,
|
||||
*,
|
||||
instant: datetime,
|
||||
consumption_kwh: float | None = 50.0,
|
||||
source: str = "api_history",
|
||||
is_working_hours: bool | None = True,
|
||||
) -> None:
|
||||
"""Une lecture isolee, quand le test pilote sa valeur plutot que sa forme."""
|
||||
connexion.execute(
|
||||
_INSERT_READING,
|
||||
{
|
||||
"site_id": site_id,
|
||||
"timestamp": instant,
|
||||
"source": source,
|
||||
"consumption_kwh": consumption_kwh,
|
||||
"temperature_celsius": 15.0,
|
||||
"humidity_percent": 50.0,
|
||||
"solar_irradiance_wm2": 0.0,
|
||||
"is_working_hours": is_working_hours,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def insere_prediction(
|
||||
connexion: Connection,
|
||||
site_id: str,
|
||||
*,
|
||||
target_at: datetime,
|
||||
predicted_value: float | None = 42.0,
|
||||
model_reference: str = "lightgbm-test000000",
|
||||
status: str = "available",
|
||||
failure_reason: str | None = None,
|
||||
) -> None:
|
||||
connexion.execute(
|
||||
_INSERT_PREDICTION,
|
||||
{
|
||||
"site_id": site_id,
|
||||
"target_at": target_at,
|
||||
"predicted_value": predicted_value,
|
||||
"model_reference": model_reference,
|
||||
"status": status,
|
||||
"failure_reason": failure_reason,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def url_ml() -> URL:
|
||||
valeur = os.environ.get("ML_DATABASE_URL")
|
||||
if not valeur:
|
||||
pytest.fail("ML_DATABASE_URL absente. Voir `make ml-test-integration`.")
|
||||
|
||||
url = make_url(valeur)
|
||||
if url.database != BASE_ATTENDUE:
|
||||
pytest.fail(
|
||||
f"Ces tests ecrivent et suppriment : ML_DATABASE_URL doit viser {BASE_ATTENDUE}, "
|
||||
f"pas {url.database}."
|
||||
)
|
||||
return url
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def moteur_ml(url_ml: URL) -> Iterator[Engine]:
|
||||
moteur = create_engine(url_ml)
|
||||
try:
|
||||
yield moteur
|
||||
finally:
|
||||
moteur.dispose()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def connexion_ml(moteur_ml: Engine) -> Iterator[Connection]:
|
||||
with moteur_ml.connect() as connexion:
|
||||
transaction = connexion.begin()
|
||||
try:
|
||||
yield connexion
|
||||
finally:
|
||||
transaction.rollback()
|
||||
|
||||
|
||||
@dataclass
|
||||
class Parc:
|
||||
"""Semis valide en base, et son nettoyage, pour les tests qui appellent `run_scoring`.
|
||||
|
||||
Chaque `site_id` porte une marque unique : la base de test est partagee entre un run local
|
||||
et la CI.
|
||||
"""
|
||||
|
||||
moteur: Engine
|
||||
sites: list[str] = field(default_factory=list)
|
||||
|
||||
def site(self, *, site_type: str = SITE_TYPE, capacity_kw: float | None = CAPACITY_KW) -> str:
|
||||
with self.moteur.begin() as connexion:
|
||||
site_id = insere_site(connexion, site_type=site_type, capacity_kw=capacity_kw)
|
||||
self.sites.append(site_id)
|
||||
return site_id
|
||||
|
||||
def lectures(self, site_id: str, **arguments: Any) -> list[datetime]:
|
||||
with self.moteur.begin() as connexion:
|
||||
return insere_lectures(connexion, site_id, **arguments)
|
||||
|
||||
def lecture(self, site_id: str, **arguments: Any) -> None:
|
||||
with self.moteur.begin() as connexion:
|
||||
insere_lecture(connexion, site_id, **arguments)
|
||||
|
||||
def prediction(self, site_id: str, **arguments: Any) -> None:
|
||||
with self.moteur.begin() as connexion:
|
||||
insere_prediction(connexion, site_id, **arguments)
|
||||
|
||||
def predictions_ecrites(self, site_id: str) -> list[Row[Any]]:
|
||||
with self.moteur.connect() as connexion:
|
||||
return list(connexion.execute(_SELECT_PREDICTIONS, {"site_id": site_id}))
|
||||
|
||||
def nettoie(self) -> None:
|
||||
if not self.sites:
|
||||
return
|
||||
|
||||
with self.moteur.begin() as connexion:
|
||||
for suppression in _SUPPRESSIONS:
|
||||
connexion.execute(suppression, {"sites": self.sites})
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def parc(moteur_ml: Engine) -> Iterator[Parc]:
|
||||
semis = Parc(moteur=moteur_ml)
|
||||
try:
|
||||
yield semis
|
||||
finally:
|
||||
semis.nettoie()
|
||||
|
||||
|
||||
def trame_synthetique(*, sites: int = 2, heures: int = 400) -> pd.DataFrame:
|
||||
"""Lectures horaires deterministes, assez longues pour que le lag de 168 h existe."""
|
||||
depart = datetime(2024, 1, 1, tzinfo=UTC)
|
||||
morceaux = [
|
||||
pd.DataFrame(
|
||||
{
|
||||
"site_id": f"SITE{numero:03d}",
|
||||
"timestamp": [depart + timedelta(hours=rang) for rang in range(heures)],
|
||||
TARGET_COLUMN: [
|
||||
50.0 + 10.0 * math.sin(rang / 12.0) + numero * 5.0 for rang in range(heures)
|
||||
],
|
||||
"temperature_celsius": 15.0,
|
||||
"humidity_percent": 50.0,
|
||||
"solar_irradiance_wm2": 0.0,
|
||||
"is_working_hours": True,
|
||||
"site_type": SITE_TYPE,
|
||||
"capacity_kw": CAPACITY_KW,
|
||||
}
|
||||
)
|
||||
for numero in range(sites)
|
||||
]
|
||||
return pd.concat(morceaux, ignore_index=True)
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def modele_jetable(tmp_path_factory: pytest.TempPathFactory) -> Path:
|
||||
"""Booster reel entraine sur une trame synthetique, ecrit dans un repertoire temporaire.
|
||||
|
||||
Ni `ml/models/` (ignore par git, et le polluer serait un effet de bord), ni
|
||||
`enervision_ml.train.train()` (qui journalise dans MLflow sans garde). Le typage `category`
|
||||
de `site_type` reproduit celui de l'entrainement : c'est le `pandas_categorical` enregistre
|
||||
dans le modele que `score()` devra retrouver.
|
||||
"""
|
||||
features = build_features(trame_synthetique()).dropna(subset=feature_columns())
|
||||
typee = features.copy()
|
||||
typee["site_type"] = typee["site_type"].astype("category")
|
||||
|
||||
donnees = lgb.Dataset(
|
||||
typee[feature_columns()],
|
||||
label=typee[TARGET_COLUMN],
|
||||
categorical_feature=["site_type"],
|
||||
)
|
||||
booster = lgb.train(
|
||||
{"objective": "regression", "num_leaves": 7, "min_data_in_leaf": 5, "verbosity": -1},
|
||||
donnees,
|
||||
num_boost_round=5,
|
||||
)
|
||||
|
||||
chemin = tmp_path_factory.mktemp("modele") / "lightgbm-consumption.txt"
|
||||
booster.save_model(str(chemin))
|
||||
return chemin
|
||||
@@ -0,0 +1,154 @@
|
||||
from datetime import timedelta
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from sqlalchemy import Connection
|
||||
|
||||
from enervision_ml.data import (
|
||||
OUTPUT_COLUMNS,
|
||||
load_from_csv,
|
||||
load_from_database,
|
||||
load_recent_from_database,
|
||||
)
|
||||
from tests.conftest import ANCRAGE, insere_lecture, insere_lectures, insere_site
|
||||
|
||||
pytestmark = pytest.mark.integration
|
||||
|
||||
|
||||
def test_load_from_database_returns_the_nine_contract_columns(connexion_ml: Connection) -> None:
|
||||
site_id = insere_site(connexion_ml)
|
||||
insere_lectures(connexion_ml, site_id, heures=3, fin=ANCRAGE)
|
||||
|
||||
frame = load_from_database(connexion_ml)
|
||||
|
||||
assert list(frame.columns) == OUTPUT_COLUMNS
|
||||
|
||||
|
||||
def test_load_from_database_joins_the_site_attributes_to_every_reading(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
site_id = insere_site(connexion_ml, site_type="factory", capacity_kw=250.0)
|
||||
insere_lectures(connexion_ml, site_id, heures=3, fin=ANCRAGE)
|
||||
|
||||
frame = load_from_database(connexion_ml)
|
||||
|
||||
mien = frame[frame["site_id"] == site_id]
|
||||
assert len(mien) == 3
|
||||
assert set(mien["site_type"]) == {"factory"}
|
||||
assert set(mien["capacity_kw"]) == {250.0}
|
||||
|
||||
|
||||
def test_load_recent_from_database_excludes_readings_before_the_since_bound(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
site_id = insere_site(connexion_ml)
|
||||
insere_lectures(connexion_ml, site_id, heures=5, fin=ANCRAGE)
|
||||
|
||||
frame = load_recent_from_database(connexion_ml, since=ANCRAGE - timedelta(hours=2))
|
||||
|
||||
assert list(frame["timestamp"]) == [
|
||||
ANCRAGE - timedelta(hours=2),
|
||||
ANCRAGE - timedelta(hours=1),
|
||||
ANCRAGE,
|
||||
]
|
||||
|
||||
|
||||
def test_load_recent_from_database_includes_a_reading_exactly_at_the_since_bound(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
site_id = insere_site(connexion_ml)
|
||||
insere_lecture(connexion_ml, site_id, instant=ANCRAGE)
|
||||
|
||||
frame = load_recent_from_database(connexion_ml, since=ANCRAGE)
|
||||
|
||||
assert len(frame) == 1
|
||||
|
||||
|
||||
def test_load_recent_from_database_keeps_timestamps_timezone_aware(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
site_id = insere_site(connexion_ml)
|
||||
insere_lecture(connexion_ml, site_id, instant=ANCRAGE)
|
||||
|
||||
frame = load_recent_from_database(connexion_ml, since=ANCRAGE)
|
||||
|
||||
assert frame["timestamp"].dt.tz is not None
|
||||
|
||||
|
||||
def test_load_recent_from_database_orders_readings_by_site_then_timestamp(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
site_id = insere_site(connexion_ml)
|
||||
for decalage in (2, 0, 1):
|
||||
insere_lecture(connexion_ml, site_id, instant=ANCRAGE + timedelta(hours=decalage))
|
||||
|
||||
frame = load_recent_from_database(connexion_ml, since=ANCRAGE)
|
||||
|
||||
assert list(frame["timestamp"]) == [
|
||||
ANCRAGE,
|
||||
ANCRAGE + timedelta(hours=1),
|
||||
ANCRAGE + timedelta(hours=2),
|
||||
]
|
||||
|
||||
|
||||
def test_load_recent_from_database_returns_the_contract_columns_even_without_any_row(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
frame = load_recent_from_database(connexion_ml, since=ANCRAGE + timedelta(days=365))
|
||||
|
||||
assert frame.empty
|
||||
assert list(frame.columns) == OUTPUT_COLUMNS
|
||||
|
||||
|
||||
def test_load_recent_from_database_types_a_fully_null_capacity_kw_as_float64(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
site_id = insere_site(connexion_ml, capacity_kw=None)
|
||||
insere_lectures(connexion_ml, site_id, heures=3, fin=ANCRAGE)
|
||||
|
||||
frame = load_recent_from_database(connexion_ml, since=ANCRAGE - timedelta(hours=2))
|
||||
|
||||
assert frame["capacity_kw"].dtype == "float64"
|
||||
assert frame["capacity_kw"].isna().all()
|
||||
|
||||
|
||||
def test_load_recent_from_database_types_a_null_is_working_hours_as_float64(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
site_id = insere_site(connexion_ml)
|
||||
insere_lecture(connexion_ml, site_id, instant=ANCRAGE, is_working_hours=None)
|
||||
insere_lecture(
|
||||
connexion_ml, site_id, instant=ANCRAGE + timedelta(hours=1), is_working_hours=True
|
||||
)
|
||||
|
||||
frame = load_recent_from_database(connexion_ml, since=ANCRAGE)
|
||||
|
||||
assert frame["is_working_hours"].dtype == "float64"
|
||||
assert list(frame["is_working_hours"].isna()) == [True, False]
|
||||
|
||||
|
||||
def test_both_loaders_produce_the_same_columns_in_the_same_order(
|
||||
connexion_ml: Connection, tmp_path: Path
|
||||
) -> None:
|
||||
site_id = insere_site(connexion_ml)
|
||||
insere_lectures(connexion_ml, site_id, heures=2, fin=ANCRAGE)
|
||||
csv_path = tmp_path / "lectures.csv"
|
||||
pd.DataFrame(
|
||||
{
|
||||
"site_id": [site_id],
|
||||
"timestamp": [ANCRAGE],
|
||||
"consumption_kwh": [50.0],
|
||||
"temperature_celsius": [15.0],
|
||||
"humidity_percent": [50.0],
|
||||
"solar_irradiance_wm2": [0.0],
|
||||
"is_working_hours": [True],
|
||||
"site_type": ["office"],
|
||||
}
|
||||
).to_csv(csv_path, index=False)
|
||||
|
||||
depuis_la_base = load_recent_from_database(connexion_ml, since=ANCRAGE - timedelta(hours=1))
|
||||
depuis_le_csv = load_from_csv(csv_path)
|
||||
|
||||
assert list(depuis_la_base.columns) == list(depuis_le_csv.columns)
|
||||
assert depuis_la_base.dtypes.to_dict() == depuis_le_csv.dtypes.to_dict()
|
||||
@@ -0,0 +1,225 @@
|
||||
from datetime import datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import Connection, Row, text
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
from enervision_ml.score import (
|
||||
INSUFFICIENT_DATA_REASON,
|
||||
LOOKBACK,
|
||||
MAX_STALENESS,
|
||||
ScoredSite,
|
||||
model_reference,
|
||||
run_scoring,
|
||||
write_predictions,
|
||||
)
|
||||
from tests.conftest import ANCRAGE, Parc, insere_site
|
||||
|
||||
pytestmark = pytest.mark.integration
|
||||
|
||||
REFERENCE = "lightgbm-000000000000"
|
||||
|
||||
_SELECT = text(
|
||||
"""
|
||||
SELECT target_at, target_metric, period_minutes, predicted_value,
|
||||
model_reference, status, failure_reason
|
||||
FROM prediction
|
||||
WHERE site_id = :site_id
|
||||
ORDER BY prediction_id
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def lignes(connexion: Connection, site_id: str) -> list[Row[Any]]:
|
||||
return list(connexion.execute(_SELECT, {"site_id": site_id}))
|
||||
|
||||
|
||||
def disponible(
|
||||
site_id: str,
|
||||
*,
|
||||
target_at: datetime = ANCRAGE,
|
||||
predicted_value: float | None = 12.5,
|
||||
) -> ScoredSite:
|
||||
return ScoredSite(
|
||||
site_id=site_id,
|
||||
target_at=target_at,
|
||||
status="available",
|
||||
predicted_value=predicted_value,
|
||||
failure_reason=None,
|
||||
)
|
||||
|
||||
|
||||
def test_write_predictions_inserts_one_row_per_scored_site(connexion_ml: Connection) -> None:
|
||||
premier = insere_site(connexion_ml)
|
||||
second = insere_site(connexion_ml)
|
||||
|
||||
write_predictions(connexion_ml, [disponible(premier), disponible(second)], reference=REFERENCE)
|
||||
|
||||
assert len(lignes(connexion_ml, premier)) == 1
|
||||
assert len(lignes(connexion_ml, second)) == 1
|
||||
|
||||
|
||||
def test_write_predictions_stores_the_model_reference_and_the_hourly_period(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
site_id = insere_site(connexion_ml)
|
||||
|
||||
write_predictions(connexion_ml, [disponible(site_id)], reference=REFERENCE)
|
||||
|
||||
ligne = lignes(connexion_ml, site_id)[0]
|
||||
assert ligne.model_reference == REFERENCE
|
||||
assert ligne.target_metric == "consumption_kwh"
|
||||
assert ligne.period_minutes == 60
|
||||
|
||||
|
||||
def test_write_predictions_stacks_a_second_run_instead_of_overwriting_the_first(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
site_id = insere_site(connexion_ml)
|
||||
|
||||
write_predictions(
|
||||
connexion_ml, [disponible(site_id, predicted_value=10.0)], reference=REFERENCE
|
||||
)
|
||||
write_predictions(
|
||||
connexion_ml, [disponible(site_id, predicted_value=20.0)], reference=REFERENCE
|
||||
)
|
||||
|
||||
assert [ligne.predicted_value for ligne in lignes(connexion_ml, site_id)] == [10.0, 20.0]
|
||||
|
||||
|
||||
def test_write_predictions_writes_nothing_when_no_site_was_scored(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
site_id = insere_site(connexion_ml)
|
||||
|
||||
write_predictions(connexion_ml, [], reference=REFERENCE)
|
||||
|
||||
assert lignes(connexion_ml, site_id) == []
|
||||
|
||||
|
||||
def test_write_predictions_rejects_an_available_row_without_a_predicted_value(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
site_id = insere_site(connexion_ml)
|
||||
|
||||
with pytest.raises(IntegrityError, match="ck_prediction_status"):
|
||||
write_predictions(
|
||||
connexion_ml, [disponible(site_id, predicted_value=None)], reference=REFERENCE
|
||||
)
|
||||
|
||||
|
||||
def test_write_predictions_rejects_an_insufficient_data_row_carrying_a_value(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
site_id = insere_site(connexion_ml)
|
||||
incoherent = ScoredSite(
|
||||
site_id=site_id,
|
||||
target_at=ANCRAGE,
|
||||
status="insufficient_data",
|
||||
predicted_value=12.5,
|
||||
failure_reason=INSUFFICIENT_DATA_REASON,
|
||||
)
|
||||
|
||||
with pytest.raises(IntegrityError, match="ck_prediction_status"):
|
||||
write_predictions(connexion_ml, [incoherent], reference=REFERENCE)
|
||||
|
||||
|
||||
def test_write_predictions_rejects_a_prediction_for_an_unknown_site(
|
||||
connexion_ml: Connection,
|
||||
) -> None:
|
||||
with pytest.raises(IntegrityError, match="fk_prediction_site"):
|
||||
write_predictions(connexion_ml, [disponible("SITE-INCONNU")], reference=REFERENCE)
|
||||
|
||||
|
||||
def test_run_scoring_writes_an_available_prediction_for_a_site_with_a_full_week(
|
||||
parc: Parc, modele_jetable: Path
|
||||
) -> None:
|
||||
site_id = parc.site()
|
||||
parc.lectures(site_id, heures=200, fin=ANCRAGE)
|
||||
|
||||
run_scoring(model_path=modele_jetable, now=ANCRAGE)
|
||||
|
||||
ligne = parc.predictions_ecrites(site_id)[0]
|
||||
assert ligne.status == "available"
|
||||
assert ligne.predicted_value is not None
|
||||
assert ligne.target_at == ANCRAGE + timedelta(hours=1)
|
||||
|
||||
|
||||
def test_run_scoring_writes_insufficient_data_when_the_weekly_lag_is_missing(
|
||||
parc: Parc, modele_jetable: Path
|
||||
) -> None:
|
||||
site_id = parc.site()
|
||||
parc.lectures(site_id, heures=100, fin=ANCRAGE)
|
||||
|
||||
run_scoring(model_path=modele_jetable, now=ANCRAGE)
|
||||
|
||||
ligne = parc.predictions_ecrites(site_id)[0]
|
||||
assert ligne.status == "insufficient_data"
|
||||
assert ligne.predicted_value is None
|
||||
assert ligne.failure_reason == INSUFFICIENT_DATA_REASON
|
||||
|
||||
|
||||
def test_run_scoring_writes_a_staleness_reason_when_the_last_reading_is_too_old(
|
||||
parc: Parc, modele_jetable: Path
|
||||
) -> None:
|
||||
site_id = parc.site()
|
||||
parc.lectures(site_id, heures=200, fin=ANCRAGE)
|
||||
|
||||
run_scoring(model_path=modele_jetable, now=ANCRAGE + MAX_STALENESS + timedelta(hours=1))
|
||||
|
||||
ligne = parc.predictions_ecrites(site_id)[0]
|
||||
assert ligne.status == "insufficient_data"
|
||||
assert ligne.failure_reason != INSUFFICIENT_DATA_REASON
|
||||
|
||||
|
||||
def test_run_scoring_writes_nothing_when_every_reading_is_older_than_the_window(
|
||||
parc: Parc, modele_jetable: Path
|
||||
) -> None:
|
||||
site_id = parc.site()
|
||||
parc.lectures(site_id, heures=200, fin=ANCRAGE)
|
||||
|
||||
run_scoring(model_path=modele_jetable, now=ANCRAGE + LOOKBACK + timedelta(days=1))
|
||||
|
||||
assert parc.predictions_ecrites(site_id) == []
|
||||
|
||||
|
||||
def test_run_scoring_only_writes_the_site_that_was_requested(
|
||||
parc: Parc, modele_jetable: Path
|
||||
) -> None:
|
||||
demande = parc.site()
|
||||
ignore = parc.site()
|
||||
parc.lectures(demande, heures=200, fin=ANCRAGE)
|
||||
parc.lectures(ignore, heures=200, fin=ANCRAGE)
|
||||
|
||||
run_scoring(model_path=modele_jetable, site_id=demande, now=ANCRAGE)
|
||||
|
||||
assert len(parc.predictions_ecrites(demande)) == 1
|
||||
assert parc.predictions_ecrites(ignore) == []
|
||||
|
||||
|
||||
def test_run_scoring_uses_the_model_file_hash_as_model_reference(
|
||||
parc: Parc, modele_jetable: Path
|
||||
) -> None:
|
||||
site_id = parc.site()
|
||||
parc.lectures(site_id, heures=200, fin=ANCRAGE)
|
||||
|
||||
run_scoring(model_path=modele_jetable, now=ANCRAGE)
|
||||
|
||||
ligne = parc.predictions_ecrites(site_id)[0]
|
||||
assert ligne.model_reference == model_reference(modele_jetable)
|
||||
|
||||
|
||||
def test_run_scoring_appends_a_second_row_when_it_runs_twice(
|
||||
parc: Parc, modele_jetable: Path
|
||||
) -> None:
|
||||
site_id = parc.site()
|
||||
parc.lectures(site_id, heures=200, fin=ANCRAGE)
|
||||
|
||||
run_scoring(model_path=modele_jetable, now=ANCRAGE)
|
||||
run_scoring(model_path=modele_jetable, now=ANCRAGE)
|
||||
|
||||
ecrites = parc.predictions_ecrites(site_id)
|
||||
assert len(ecrites) == 2
|
||||
assert ecrites[0].target_at == ecrites[1].target_at
|
||||
Reference in New Issue
Block a user