Spaces:
Running
Running
File size: 4,341 Bytes
a721dfa 9734b71 a721dfa 9734b71 a721dfa 9734b71 a721dfa 9734b71 a721dfa 9734b71 a721dfa 9734b71 a721dfa 9734b71 a721dfa | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 | 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()
|