SuperAI_Forecast / backend /test_startup_utils.py
Thang6822
Update branding to SuperAI Forecast
9734b71
Raw
History Blame
4.34 kB
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()