sdpy-kit 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
kit/ai/pool.py ADDED
@@ -0,0 +1,411 @@
1
+ """
2
+ Copyright (c) 2021-2026 Clark Chang. All Rights Reserved.
3
+
4
+ Licensed under the Apache License, Version 2.0 (the "License");
5
+ you may not use this file except in compliance with the License.
6
+ You may obtain a copy of the License at
7
+
8
+ http://www.apache.org/licenses/LICENSE-2.0
9
+
10
+ Unless required by applicable law or agreed to in writing, software
11
+ distributed under the License is distributed on an "AS IS" BASIS,
12
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13
+ See the License for the specific language governing permissions and
14
+ limitations under the License.
15
+
16
+ Project: SynthArk
17
+ Author: Clark Chang
18
+ """
19
+
20
+ # date: 2026-09-05
21
+ """AI 服务候选池:激活探测、故障冷却、当前项切换
22
+
23
+ 按模型类型维护「候选链」:链上档案按优先级排序(``prefer()`` 指定项打头),
24
+ 仅 ``AVAILABLE`` 状态档案参与调用竞选;档案调用失败进入冷却(默认 60 秒),
25
+ 冷却期内被跳过(备胎顶上),到期自动恢复——冷却即粘性降级,自愈无需干预。
26
+
27
+ 预热探测(``probe_all``)经 ``GET {base_url}/models`` 验证密钥有效性:
28
+ 2xx / 429 视为 key 有效并激活;401 / 403 判定 key 无效出链;其余
29
+ (超时 / 404 / 5xx)无法判定出链。剔除后链空时回退「密钥已配置档案全部
30
+ 惰性入链」并告警,保证功能不瘫。
31
+
32
+ 线程模型:业务调用与探测均在 runner 后台事件循环线程内执行;
33
+ ``available()`` / ``prefer()`` 可能从业务线程调用,内部以锁保护共享状态。
34
+ """
35
+
36
+ import asyncio
37
+ import threading
38
+ import time
39
+ from collections.abc import Callable, Iterator
40
+ from dataclasses import dataclass, field
41
+ from typing import Final
42
+
43
+ from kit.ai.chat_service import ChatService
44
+ from kit.ai.embedding_service import EmbeddingService
45
+ from kit.ai.exceptions import AIConfigError, AIError, AIRequestError
46
+ from kit.ai.image_service import ImageService
47
+ from kit.ai.model_builder import Service
48
+ from kit.ai.model_client import ModelClient
49
+ from kit.ai.profiles import (
50
+ ModelProfile,
51
+ ModelRegistry,
52
+ ProfileState,
53
+ ProfileStatus,
54
+ get_model_registry,
55
+ )
56
+ from kit.ai.rerank_service import RerankService
57
+ from kit.config.log_config import get_default_logger
58
+
59
+ __all__ = ["AIServicePool"]
60
+
61
+ logger = get_default_logger(__name__)
62
+
63
+ # 调用失败后的冷却时长(秒):冷却期内备胎顶上,到期自动回位
64
+ _COOLDOWN_SECONDS: Final = 60.0
65
+
66
+ # 各模型类型 → 单档案服务工厂
67
+ _FACTORIES: Final[dict[str, Callable[[ModelProfile], Service]]] = {
68
+ "chat": ChatService,
69
+ "embedding": EmbeddingService,
70
+ "rerank": RerankService,
71
+ "image": ImageService,
72
+ }
73
+
74
+ # 探测判定为「密钥无效」的 HTTP 状态码(2xx 已由 list_models 放行;
75
+ # 429 限流恰证明服务与密钥均有效,在 _probe_state 中特判为可用)
76
+ _PROBE_INVALID_STATUS: Final[frozenset[int]] = frozenset({401, 403})
77
+
78
+
79
+ async def _probe_profile(profile: ModelProfile) -> None:
80
+ """单档案探测:key 有效(2xx / 429)静默返回,否则抛 ``AIError``。
81
+
82
+ 独立模块级函数便于测试替换。探测客户端用毕即关。
83
+ """
84
+ client = ModelClient(profile)
85
+ try:
86
+ await client.list_models()
87
+ finally:
88
+ await client.close()
89
+
90
+
91
+ def _probe_state(exc: BaseException) -> ProfileState:
92
+ """探测异常 → 档案状态:429 视为 key 有效;401/403 无效;其余无法判定。"""
93
+ if isinstance(exc, AIRequestError) and exc.status_code == 429:
94
+ return ProfileState.AVAILABLE
95
+ if (
96
+ isinstance(exc, AIRequestError)
97
+ and exc.status_code in _PROBE_INVALID_STATUS
98
+ ):
99
+ return ProfileState.INVALID_KEY
100
+ return ProfileState.UNREACHABLE
101
+
102
+
103
+ class _Entry:
104
+ """链内单档案的运行时状态(可变)。"""
105
+
106
+ __slots__ = (
107
+ "profile",
108
+ "state",
109
+ "cooling_until",
110
+ "last_error",
111
+ "last_ok_at",
112
+ "service",
113
+ )
114
+
115
+ def __init__(self, profile: ModelProfile) -> None:
116
+ self.profile: ModelProfile = profile
117
+ self.state: ProfileState = ProfileState.AVAILABLE
118
+ # 冷却截止时刻(time.monotonic 秒);0 表示未冷却
119
+ self.cooling_until: float = 0.0
120
+ self.last_error: str | None = None
121
+ self.last_ok_at: float | None = None
122
+ self.service: Service | None = None
123
+
124
+
125
+ @dataclass
126
+ class _Chain:
127
+ """某类型的候选链:链序 + 全档案状态表(entries 含被剔除档案)。"""
128
+
129
+ order: list[str] = field(default_factory=list)
130
+ entries: dict[str, _Entry] = field(default_factory=dict)
131
+
132
+
133
+ class AIServicePool:
134
+ """按模型类型维护候选链:状态表 + 冷却 + 当前项 + service 惰性单例。"""
135
+
136
+ def __init__(self, registry: ModelRegistry | None = None) -> None:
137
+ self._registry: ModelRegistry = registry or get_model_registry()
138
+ self._lock = threading.Lock()
139
+ self._chains: dict[str, _Chain] = {}
140
+ # prefer() 指定的长期当前项(类型 → 档案名);进程重启即失效
141
+ self._preferred: dict[str, str] = {}
142
+
143
+ # ---------- 链解析 ----------
144
+
145
+ def _ensure_chain(self, model_type: str) -> _Chain:
146
+ """取(或首次解析)类型候选链。
147
+
148
+ ``registry.chain`` 涉及读配置与环境变量,置于锁外执行;
149
+ 锁内双检避免并发重复解析。
150
+ """
151
+ with self._lock:
152
+ chain = self._chains.get(model_type)
153
+ if chain is not None:
154
+ return chain
155
+ profiles = self._registry.chain(model_type)
156
+ with self._lock:
157
+ chain = self._chains.get(model_type)
158
+ if chain is not None:
159
+ return chain
160
+ entries: dict[str, _Entry] = {}
161
+ order: list[str] = []
162
+ for profile in profiles:
163
+ if profile.name not in entries:
164
+ entries[profile.name] = _Entry(profile)
165
+ order.append(profile.name)
166
+ # embedding 链内强制同 model:维度一致性是向量库可用性的前提
167
+ if model_type == "embedding" and order:
168
+ base_model = entries[order[0]].profile.model
169
+ excluded = [
170
+ name
171
+ for name in order
172
+ if entries[name].profile.model != base_model
173
+ ]
174
+ for name in excluded:
175
+ entries[name].state = ProfileState.EXCLUDED
176
+ entries[name].last_error = (
177
+ f"embedding 链内 model 不一致(要求 {base_model})"
178
+ )
179
+ order.remove(name)
180
+ if excluded:
181
+ logger.warning(
182
+ "embedding 候选链剔除 model 不一致档案: "
183
+ f"{', '.join(excluded)}(链内仅允许 {base_model})"
184
+ )
185
+ chain = _Chain(order=order, entries=entries)
186
+ self._chains[model_type] = chain
187
+ return chain
188
+
189
+ # ---------- 探测 ----------
190
+
191
+ async def probe_all(self) -> None:
192
+ """后台预热探测:逐类型并行验证候选档案密钥有效性(不抛异常)。
193
+
194
+ 单类型未配置(json 缺该类型)仅跳过,不影响其余类型探测。
195
+ """
196
+ for model_type in _FACTORIES:
197
+ try:
198
+ await self._probe_type(model_type)
199
+ except AIConfigError as exc:
200
+ logger.debug(f"{model_type} 探测跳过(类型未配置档案): {exc}")
201
+ except Exception as exc: # noqa: BLE001 - 预热失败不影响主流程
202
+ logger.warning(f"{model_type} 档案探测异常: {exc}")
203
+
204
+ async def _probe_type(self, model_type: str) -> None:
205
+ """探测单类型链上档案,按判定表激活或剔除;链空时回退惰性入链。"""
206
+ chain = self._ensure_chain(model_type)
207
+ with self._lock:
208
+ names = list(chain.order)
209
+ if not names:
210
+ return
211
+ results = await asyncio.gather(
212
+ *(_probe_profile(chain.entries[name].profile) for name in names),
213
+ return_exceptions=True,
214
+ )
215
+ dropped: list[str] = []
216
+ with self._lock:
217
+ for name, result in zip(names, results, strict=True):
218
+ entry = chain.entries[name]
219
+ if isinstance(result, BaseException):
220
+ state = _probe_state(result)
221
+ entry.state = state
222
+ entry.last_error = f"探测: {result}"
223
+ if state is ProfileState.AVAILABLE:
224
+ # 429:限流但 key 有效,保持可用
225
+ entry.last_ok_at = time.time()
226
+ else:
227
+ dropped.append(name)
228
+ logger.warning(
229
+ f"{model_type}.{name} 探测未通过,剔除出链: {result}"
230
+ )
231
+ else:
232
+ entry.state = ProfileState.AVAILABLE
233
+ entry.last_ok_at = time.time()
234
+ entry.last_error = None
235
+ if dropped:
236
+ for name in dropped:
237
+ if name in chain.order:
238
+ chain.order.remove(name)
239
+ # 兜底:剔除后链空 → 回退「密钥已配置档案全部惰性入链」
240
+ if not chain.order:
241
+ chain.order = list(names)
242
+ for name in names:
243
+ chain.entries[name].state = ProfileState.AVAILABLE
244
+ logger.warning(
245
+ f"{model_type} 探测后候选链为空,回退惰性入链: "
246
+ f"{', '.join(names)}(真实调用时再验证)"
247
+ )
248
+ else:
249
+ logger.info(
250
+ f"{model_type} 探测完成: 可用 "
251
+ f"{', '.join(chain.order)};剔除 {', '.join(dropped)}"
252
+ )
253
+
254
+ # ---------- 候选迭代 ----------
255
+
256
+ def iter_services(self, model_type: str) -> Iterator[tuple[str, Service]]:
257
+ """按竞选顺序迭代(档案名, service):prefer 项打头,跳过冷却档案。
258
+
259
+ 候选计划(顺序 + prefer 提前 + 冷却到期复位 + 可用性过滤)在生成器
260
+ 首次推进时一次定格:此后 ``prefer()`` 切换与其它请求的失败冷却均
261
+ 不影响已开始迭代的调用方——切换只为「下一次」请求生效。
262
+
263
+ 仅在后台事件循环线程内迭代;service 惰性构建并缓存复用。
264
+ """
265
+ chain = self._ensure_chain(model_type)
266
+ now = time.monotonic()
267
+ with self._lock:
268
+ seq = list(chain.order)
269
+ preferred = self._preferred.get(model_type)
270
+ if preferred is not None and preferred in seq:
271
+ seq.remove(preferred)
272
+ seq.insert(0, preferred)
273
+ plan: list[tuple[str, Service | None]] = []
274
+ for name in seq:
275
+ entry = chain.entries[name]
276
+ if (
277
+ entry.state is ProfileState.COOLING
278
+ and now >= entry.cooling_until
279
+ ):
280
+ entry.state = ProfileState.AVAILABLE
281
+ entry.last_error = None
282
+ if entry.state is not ProfileState.AVAILABLE:
283
+ continue
284
+ plan.append((name, entry.service))
285
+ for name, service in plan:
286
+ if service is None:
287
+ entry = chain.entries[name]
288
+ service = _FACTORIES[model_type](entry.profile)
289
+ with self._lock:
290
+ if entry.service is None:
291
+ entry.service = service
292
+ else:
293
+ service = entry.service
294
+ yield name, service
295
+
296
+ # ---------- 状态标记 ----------
297
+
298
+ def mark_failure(self, model_type: str, name: str, exc: AIError) -> None:
299
+ """档案调用失败:进入冷却并记录原因。"""
300
+ with self._lock:
301
+ chain = self._chains.get(model_type)
302
+ entry = chain.entries.get(name) if chain else None
303
+ if entry is None:
304
+ return
305
+ entry.state = ProfileState.COOLING
306
+ entry.cooling_until = time.monotonic() + _COOLDOWN_SECONDS
307
+ entry.last_error = f"{type(exc).__name__}: {exc}"
308
+ logger.warning(
309
+ f"{model_type}.{name} 调用失败,冷却 {_COOLDOWN_SECONDS:g}s: {exc}"
310
+ )
311
+
312
+ def mark_success(self, model_type: str, name: str) -> None:
313
+ """档案调用成功:记录成功时刻并清除失败痕迹。"""
314
+ with self._lock:
315
+ chain = self._chains.get(model_type)
316
+ entry = chain.entries.get(name) if chain else None
317
+ if entry is None:
318
+ return
319
+ entry.last_ok_at = time.time()
320
+ entry.last_error = None
321
+
322
+ # ---------- 当前项与状态 ----------
323
+
324
+ def prefer(self, model_type: str, profile: str) -> None:
325
+ """切换该类型长期当前项:指定档案成为固定链头并清除冷却。
326
+
327
+ Raises:
328
+ AIConfigError: 档案不存在,或不可用
329
+ (未配置 / 密钥无效 / 不可达 / 因约束被排除)。
330
+ """
331
+ chain = self._ensure_chain(model_type)
332
+ with self._lock:
333
+ entry = chain.entries.get(profile)
334
+ if entry is None:
335
+ known = ", ".join(chain.entries) or "无"
336
+ raise AIConfigError(
337
+ f"档案不存在: {model_type}.{profile}(已知档案: {known})"
338
+ )
339
+ if entry.state in (
340
+ ProfileState.UNCONFIGURED,
341
+ ProfileState.INVALID_KEY,
342
+ ProfileState.UNREACHABLE,
343
+ ProfileState.EXCLUDED,
344
+ ):
345
+ raise AIConfigError(
346
+ f"档案不可用: {model_type}.{profile}"
347
+ f"(状态 {entry.state.value},原因: {entry.last_error})"
348
+ )
349
+ entry.state = ProfileState.AVAILABLE
350
+ entry.cooling_until = 0.0
351
+ self._preferred[model_type] = profile
352
+ logger.info(f"已切换 {model_type} 当前使用档案: {profile}")
353
+
354
+ def available(self, model_type: str) -> list[ProfileStatus]:
355
+ """该类型全部档案状态快照(含未配置与被剔除者),链头标记 current。"""
356
+ chain = self._ensure_chain(model_type)
357
+ now = time.monotonic()
358
+ with self._lock:
359
+ preferred = self._preferred.get(model_type)
360
+ if preferred in chain.order:
361
+ current: str | None = preferred
362
+ else:
363
+ current = chain.order[0] if chain.order else None
364
+ rows = [
365
+ ProfileStatus(
366
+ name=name,
367
+ model=entry.profile.model,
368
+ state=(
369
+ ProfileState.AVAILABLE
370
+ if entry.state is ProfileState.COOLING
371
+ and now >= entry.cooling_until
372
+ else entry.state
373
+ ),
374
+ current=name == current,
375
+ last_error=entry.last_error,
376
+ last_ok_at=entry.last_ok_at,
377
+ )
378
+ for name, entry in chain.entries.items()
379
+ ]
380
+ try:
381
+ all_profiles = self._registry.type_profiles(model_type)
382
+ except AIConfigError:
383
+ all_profiles = {}
384
+ for name, config in all_profiles.items():
385
+ if name not in chain.entries:
386
+ rows.append(
387
+ ProfileStatus(
388
+ name=name,
389
+ model=config.model,
390
+ state=ProfileState.UNCONFIGURED,
391
+ current=False,
392
+ )
393
+ )
394
+ return rows
395
+
396
+ # ---------- 生命周期 ----------
397
+
398
+ async def close_all(self) -> None:
399
+ """关闭全部已构建 service(释放底层 httpx 连接)。幂等。"""
400
+ with self._lock:
401
+ services = [
402
+ entry.service
403
+ for chain in self._chains.values()
404
+ for entry in chain.entries.values()
405
+ if entry.service is not None
406
+ ]
407
+ for service in services:
408
+ try:
409
+ await service.close()
410
+ except Exception as exc: # noqa: BLE001 - 关闭阶段失败不影响退出
411
+ logger.warning(f"关闭 {service} 失败: {exc}")