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()