File size: 5,282 Bytes
5f771b4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
import asyncio
import sys
import time
import types
import unittest
from enum import IntEnum
from pathlib import Path

from pydantic import ValidationError

BACKEND_ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(BACKEND_ROOT))


def _install_import_stubs() -> None:
    auth_guard = types.ModuleType("api.auth_guard")

    class AuthRole(IntEnum):
        MACHINE = 1

    auth_guard.AuthRole = AuthRole
    auth_guard.require_role = lambda _role: (lambda: None)
    sys.modules["api.auth_guard"] = auth_guard

    health_module = types.ModuleType("api.health_manager")

    class HealthyManager:
        async def is_healthy(self, _worker_id: str) -> bool:
            return True

    health_module.health_manager = HealthyManager()
    sys.modules["api.health_manager"] = health_module


_install_import_stubs()

from api.marketplace import WORKERS_REGISTRY, WorkerCapability
from api.resolver import CapabilityResolver, ResolverConstraints


class CapabilityResolverTests(unittest.IsolatedAsyncioTestCase):
    def setUp(self):
        self._original_registry = dict(WORKERS_REGISTRY)
        WORKERS_REGISTRY.clear()

    def tearDown(self):
        WORKERS_REGISTRY.clear()
        WORKERS_REGISTRY.update(self._original_registry)

    def register_worker(self, **overrides) -> WorkerCapability:
        values = {
            "id": "worker",
            "name": "Worker",
            "version": "1.0.0",
            "last_seen": int(time.time()),
            "capabilities": ["vision"],
            "cost": 1.0,
            "latency": 100.0,
            "region": "global",
            "gpu": False,
            "priority": 10,
        }
        values.update(overrides)
        worker = WorkerCapability(**values)
        WORKERS_REGISTRY[worker.id] = worker
        return worker

    async def test_uses_semver_not_lexicographic_ordering(self):
        self.register_worker(id="v2", version="2.0.0", priority=20)
        self.register_worker(id="v10", version="10.0.0", priority=10)

        worker = await CapabilityResolver.resolve(
            "vision",
            ResolverConstraints(min_version="2.0.0"),
        )

        self.assertIsNotNone(worker)
        self.assertEqual(worker.id, "v10")

    async def test_excludes_prerelease_below_stable_minimum(self):
        self.register_worker(id="candidate", version="2.0.0-rc.1", priority=1)
        self.register_worker(id="stable", version="2.0.0", priority=10)

        worker = await CapabilityResolver.resolve(
            "vision",
            ResolverConstraints(min_version="2.0.0"),
        )

        self.assertIsNotNone(worker)
        self.assertEqual(worker.id, "stable")

    def test_rejects_malformed_minimum_semver(self):
        with self.assertRaises(ValidationError):
            ResolverConstraints(min_version="2.0")

    async def test_excludes_malformed_worker_version_only_when_constrained(self):
        self.register_worker(id="malformed", version="not-a-version", priority=1)
        self.register_worker(id="valid", version="2.0.0", priority=10)

        worker = await CapabilityResolver.resolve(
            "vision",
            ResolverConstraints(min_version="2.0.0"),
        )

        self.assertIsNotNone(worker)
        self.assertEqual(worker.id, "valid")

    async def test_prefers_requested_region_when_available(self):
        self.register_worker(id="global", region="global", priority=1)
        self.register_worker(id="eu", region="eu-west", priority=20)

        worker = await CapabilityResolver.resolve(
            "vision",
            ResolverConstraints(preferred_region="eu-west"),
        )

        self.assertIsNotNone(worker)
        self.assertEqual(worker.id, "eu")

    async def test_falls_back_when_requested_region_is_unavailable(self):
        self.register_worker(id="global", region="global", priority=1)

        worker = await CapabilityResolver.resolve(
            "vision",
            ResolverConstraints(preferred_region="eu-west"),
        )

        self.assertIsNotNone(worker)
        self.assertEqual(worker.id, "global")

    async def test_applies_priority_threshold_with_lower_values_preferred(self):
        self.register_worker(id="allowed", priority=100, cost=10.0)
        self.register_worker(id="excluded", priority=101, cost=0.0)

        worker = await CapabilityResolver.resolve(
            "vision",
            ResolverConstraints(min_priority=100),
        )

        self.assertIsNotNone(worker)
        self.assertEqual(worker.id, "allowed")

    async def test_preserves_existing_cost_latency_and_gpu_constraints(self):
        self.register_worker(
            id="cpu",
            cost=1.0,
            latency=50.0,
            gpu=False,
            priority=1,
        )
        self.register_worker(
            id="gpu",
            cost=2.0,
            latency=100.0,
            gpu=True,
            priority=10,
        )

        worker = await CapabilityResolver.resolve(
            "vision",
            ResolverConstraints(
                max_cost=2.0,
                max_latency=100.0,
                require_gpu=True,
            ),
        )

        self.assertIsNotNone(worker)
        self.assertEqual(worker.id, "gpu")


if __name__ == "__main__":
    unittest.main()