from __future__ import annotations from collections import defaultdict import unittest from backend.startup_utils import ( build_source_selftest_urls, clear_stale_ip_limits, run_source_selftest, warmup_timesfm, ) class _FakeLogger: def __init__(self) -> None: self.infos: list[tuple[str, tuple[object, ...]]] = [] self.warnings: list[tuple[str, tuple[object, ...]]] = [] def info(self, msg: str, *args: object, **kwargs: object) -> None: self.infos.append((msg, args)) def warning(self, msg: str, *args: object, **kwargs: object) -> None: self.warnings.append((msg, args)) class _FakeResponse: def __init__(self, status_code: int) -> None: self.status_code = status_code class _FakeClient: def __init__(self, responses: dict[str, object]) -> None: self.responses = responses async def get(self, url: str) -> object: response = self.responses[url] if isinstance(response, Exception): raise response return response class _FakeForecaster: def __init__(self, *, should_fail: bool = False) -> None: self.is_ready = not should_fail self.device = "cpu" self.should_fail = should_fail async def _lazy_load(self) -> None: if self.should_fail: raise RuntimeError("boom") class StartupUtilsTests(unittest.IsolatedAsyncioTestCase): def test_clear_stale_ip_limits_removes_only_old_entries(self) -> None: ip_limits = defaultdict(list, { "fresh": [95.0], "stale": [1.0, 2.0], }) removed = clear_stale_ip_limits(ip_limits, now=100.0, window_seconds=10) self.assertEqual(removed, 1) self.assertIn("fresh", ip_limits) self.assertNotIn("stale", ip_limits) def test_build_source_selftest_urls_includes_configured_keys(self) -> None: urls = dict(build_source_selftest_urls("td-key", "fh-key")) self.assertIn("apikey=td-key", urls["twelvedata"]) self.assertIn("token=fh-key", urls["finnhub"]) async def test_run_source_selftest_tracks_success_and_failure(self) -> None: logger = _FakeLogger() startup_sources: dict[str, dict[str, object]] = {} tests = [ ("binance", "https://binance.test"), ("broken", "https://broken.test"), ] client = _FakeClient({ "https://binance.test": _FakeResponse(200), "https://broken.test": RuntimeError("network down"), }) await run_source_selftest( client=client, tests=tests, startup_sources=startup_sources, logger=logger, timestamp_provider=lambda: "2026-04-26T00:00:00+00:00", ) self.assertEqual(startup_sources["binance"]["status_code"], 200) self.assertTrue(startup_sources["binance"]["reachable"]) self.assertFalse(startup_sources["broken"]["reachable"]) self.assertEqual(startup_sources["broken"]["error"], "network down") self.assertEqual(startup_sources["binance"]["checked_at"], "2026-04-26T00:00:00+00:00") async def test_warmup_timesfm_updates_success_state(self) -> None: logger = _FakeLogger() state = {"warming": False, "last_error": "old", "loaded": False, "device": "not_loaded"} forecaster = _FakeForecaster() await warmup_timesfm( forecaster=forecaster, startup_timesfm_state=state, logger=logger, ) self.assertFalse(state["warming"]) self.assertIsNone(state["last_error"]) self.assertTrue(state["loaded"]) self.assertEqual(state["device"], "cpu") async def test_warmup_timesfm_tracks_failure_state(self) -> None: logger = _FakeLogger() state = {"warming": False, "last_error": None, "loaded": True, "device": "not_loaded"} forecaster = _FakeForecaster(should_fail=True) await warmup_timesfm( forecaster=forecaster, startup_timesfm_state=state, logger=logger, ) self.assertFalse(state["warming"]) self.assertFalse(state["loaded"]) self.assertEqual(state["device"], "cpu") self.assertEqual(state["last_error"], "boom") if __name__ == "__main__": unittest.main()