72 lines
4.3 KiB
Python
72 lines
4.3 KiB
Python
import ast
|
|
import contextlib
|
|
import datetime
|
|
import io
|
|
from pathlib import Path
|
|
from types import SimpleNamespace
|
|
import unittest
|
|
from unittest.mock import Mock
|
|
from model_isolation import collect_predictions
|
|
|
|
|
|
class ModelIsolationTest(unittest.TestCase):
|
|
def setUp(self):
|
|
self.index = [datetime.datetime(2026,10,2,6,0) + datetime.timedelta(minutes=5*i) for i in range(3)]
|
|
self.data = {'df_fut': SimpleNamespace(index=self.index)}
|
|
self.good = {t: 1000.0 for t in self.index}
|
|
def model(self, values=None):
|
|
return SimpleNamespace(predict=Mock(return_value=self.good if values is None else values))
|
|
def call(self, models, enabled=lambda c,n:True):
|
|
return collect_predictions(self.data, {}, models, enabled)
|
|
def test_one_failed_model_does_not_remove_other_families(self):
|
|
bad=self.model();bad.predict.side_effect=ValueError('private path not logged')
|
|
forecasts,errors=self.call({1:self.model(),10:bad,21:self.model()})
|
|
self.assertEqual(forecasts[1],self.good);self.assertEqual(forecasts[21],self.good)
|
|
self.assertEqual(forecasts[10],{});self.assertEqual(errors[10]['errorType'],'ValueError')
|
|
self.assertNotIn('private',str(errors))
|
|
def test_missing_timestamp_does_not_get_filled(self):
|
|
forecasts,errors=self.call({1:self.model(),10:self.model({self.index[0]:20.0})})
|
|
self.assertFalse(forecasts[10]);self.assertIn(10,errors)
|
|
def test_nonfinite_negative_and_boolean_rejected(self):
|
|
for value in (float('nan'),float('inf'),-1.0,True,'2'):
|
|
with self.subTest(value=value):
|
|
result,errors=self.call({1:self.model(),10:self.model({t:value for t in self.index})})
|
|
self.assertFalse(result[10]);self.assertIn(10,errors)
|
|
def test_explicit_zero_forecast_is_not_imputed(self):
|
|
result,errors=self.call({10:self.model({t:0.0 for t in self.index})})
|
|
self.assertEqual(errors,{});self.assertEqual(sum(result[10].values()),0.)
|
|
def test_all_requested_models_fail_closed(self):
|
|
with self.assertRaisesRegex(ValueError,'All requested'):
|
|
self.call({10:self.model({})})
|
|
def test_disabled_models_not_called(self):
|
|
model=self.model();result,errors=self.call({10:model},lambda c,n:False)
|
|
model.predict.assert_not_called();self.assertEqual(result,{10:{}});self.assertEqual(errors,{})
|
|
def test_input_time_order_preserved(self):
|
|
result,_=self.call({1:self.model(dict(reversed(list(self.good.items()))))})
|
|
self.assertEqual(list(result[1]),self.index)
|
|
def test_run_forecast_publishes_valid_pairs_after_repeat_failure(self):
|
|
root=Path(__file__).resolve().parents[1]
|
|
tree=ast.parse((root/'main.py').read_text())
|
|
fn=next(n for n in tree.body if isinstance(n,ast.FunctionDef) and n.name=='run_forecast')
|
|
modules={n:self.model() for n in (1,2,3,10,11,13,21,22,23)}
|
|
modules[10].predict.side_effect=ValueError('profile gap')
|
|
client=Mock();published=Mock()
|
|
data={**self.data,'current_soc_perc':20.,'current_soc_source':'telemetry'}
|
|
ns={'datetime':datetime,'LOCAL_TZ':datetime.timezone.utc,'traceback':Mock(),
|
|
'get_configs':lambda:[{'anlagen_id':'test','batt_capacity_kwh':10.}],
|
|
'InfluxDBClient':Mock(return_value=client),'SYNCHRONOUS':object(),
|
|
'INFLUX_URL':'offline','INFLUX_TOKEN':'synthetic','INFLUX_ORG':'offline',
|
|
'INFLUX_BUCKET':'offline','INFLUX_TIMEOUT_MS':1,
|
|
'build_data_object':lambda *a,**k:data,'active':lambda c,n:True,
|
|
'require_recent_telemetry':lambda *a,**k:{},'collect_predictions':collect_predictions,
|
|
'_v4_publish_forecasts':published,'battery_soc_points':lambda *a:[],
|
|
'_forecast_point':lambda *a:a,'_snapshot_point':lambda *a:a,
|
|
'write_quality_metrics':Mock(),**{'v'+str(n):m for n,m in modules.items()}}
|
|
exec(compile(ast.Module(body=[fn],type_ignores=[]),'source-run-forecast','exec'),ns)
|
|
with contextlib.redirect_stdout(io.StringIO()):r=ns['run_forecast']()
|
|
self.assertEqual(r['completed'],['test'])
|
|
families=published.call_args.args[1]
|
|
self.assertTrue(families[0][1] and families[0][2]);self.assertFalse(families[1][1])
|
|
self.assertTrue(families[2][1] and families[2][2]);modules[13].predict.assert_not_called()
|
|
client.write_api.return_value.write.assert_called_once()
|