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/__init__.py +68 -0
- kit/ai/__init__.py +47 -0
- kit/ai/base_service.py +94 -0
- kit/ai/chat_service.py +110 -0
- kit/ai/embedding_service.py +133 -0
- kit/ai/exceptions.py +57 -0
- kit/ai/failover.py +244 -0
- kit/ai/image_service.py +99 -0
- kit/ai/model_builder.py +119 -0
- kit/ai/model_client.py +214 -0
- kit/ai/pool.py +411 -0
- kit/ai/profiles.py +373 -0
- kit/ai/rerank_service.py +93 -0
- kit/ai/runner.py +580 -0
- kit/ai/types.py +28 -0
- kit/common_kit.py +221 -0
- kit/config/__init__.py +44 -0
- kit/config/kit_config.py +147 -0
- kit/config/log_config.py +134 -0
- kit/db/__init__.py +70 -0
- kit/db/base_db_kit.py +406 -0
- kit/db/crud.py +361 -0
- kit/db/engine.py +230 -0
- kit/db/exceptions.py +34 -0
- kit/db/models/__init__.py +33 -0
- kit/db/models/base.py +56 -0
- kit/db/mysql_kit.py +89 -0
- kit/db/pg_kit.py +87 -0
- kit/db/session.py +85 -0
- kit/db/sqlite_kit.py +284 -0
- kit/db/transaction.py +109 -0
- kit/doc_kit.py +1502 -0
- kit/file_kit.py +429 -0
- kit/mcp_kit.py +547 -0
- kit/mineru_kit.py +918 -0
- kit/redis_kit.py +392 -0
- kit/zvec_kit.py +671 -0
- sdpy_kit-0.1.0.dist-info/METADATA +128 -0
- sdpy_kit-0.1.0.dist-info/RECORD +40 -0
- sdpy_kit-0.1.0.dist-info/WHEEL +4 -0
kit/ai/runner.py
ADDED
|
@@ -0,0 +1,580 @@
|
|
|
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: sdpy
|
|
17
|
+
Author: Clark Chang
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
from __future__ import annotations
|
|
21
|
+
|
|
22
|
+
# date: 2026-08-24
|
|
23
|
+
"""kit/ai 同步门面(共享后台事件循环 + 常驻链式服务单例)
|
|
24
|
+
|
|
25
|
+
``kit.ai`` 各服务(ChatService / EmbeddingService / RerankService / ImageService)
|
|
26
|
+
均为全异步,且实例持有绑定事件循环的 ``httpx.AsyncClient``;而业务层(service)
|
|
27
|
+
多数调用点运行在同步上下文(如 ``asyncio.to_thread`` 线程、被 FastAPI 路由直调的
|
|
28
|
+
同步函数)。
|
|
29
|
+
|
|
30
|
+
本模块用 Python 3.13 标准库实现统一适配:
|
|
31
|
+
|
|
32
|
+
- 一个由 ``asyncio.Runner`` 驱动的**共享后台事件循环线程**(daemon),常驻复用,
|
|
33
|
+
避免每次调用反复创建线程池与事件循环。
|
|
34
|
+
- 通过 ``asyncio.run_coroutine_threadsafe`` 把协程从**同步上下文**(含线程、
|
|
35
|
+
``asyncio.to_thread`` 线程)投递到后台 loop 并同步阻塞取结果——无需判断
|
|
36
|
+
调用线程是否有事件循环,也规避了 ``asyncio.run()`` 在已运行 loop 内被禁的坑。
|
|
37
|
+
- 按模型类型**缓存常驻 Failover 链式服务单例**:内部经 ``AIServicePool``
|
|
38
|
+
按候选链顺序调用(优先级链 + 失败冷却 + 自动切换),并复用底层
|
|
39
|
+
``httpx.AsyncClient`` 连接池。所有 I/O 都发生在同一后台 loop。
|
|
40
|
+
|
|
41
|
+
用法::
|
|
42
|
+
|
|
43
|
+
from kit.ai.runner import AI, shutdown
|
|
44
|
+
text = AI.chat_complete(prompt) # 同步
|
|
45
|
+
vectors = AI.embed_texts(chunks) # 同步
|
|
46
|
+
hits = AI.rerank(query, docs, top_n) # 同步
|
|
47
|
+
url = AI.generate_image(prompt) # 同步
|
|
48
|
+
AI.use("chat", "glm") # 切换 chat 当前模型档案(长期生效)
|
|
49
|
+
print(AI.available("chat")) # 查看各档案状态快照
|
|
50
|
+
# 应用退出时
|
|
51
|
+
shutdown()
|
|
52
|
+
|
|
53
|
+
注意:
|
|
54
|
+
- 本门面仅面向**同步上下文**。若调用点本身是 async 函数(事件循环内),
|
|
55
|
+
应直接 ``await`` 原生 ``ChatService.ask/stream`` 等异步接口,
|
|
56
|
+
而非经本门面阻塞事件循环线程。
|
|
57
|
+
- ``chat_stream`` 提供同步流式迭代(每次取下一帧),适合同步线程内逐 token 消费。
|
|
58
|
+
- ``AI.use`` / ``AI.available`` 为进程内内存态(不持久化),进程重启后
|
|
59
|
+
回到 env(``AI_{TYPE}_PROFILES``)/ JSON 声明顺序。
|
|
60
|
+
"""
|
|
61
|
+
|
|
62
|
+
import asyncio
|
|
63
|
+
import concurrent.futures
|
|
64
|
+
import queue
|
|
65
|
+
import threading
|
|
66
|
+
from collections.abc import Callable, Coroutine, Iterator
|
|
67
|
+
from typing import Any, Final, TypeVar, cast
|
|
68
|
+
|
|
69
|
+
from kit.ai.exceptions import AIConfigError
|
|
70
|
+
from kit.ai.failover import (
|
|
71
|
+
FailoverChatService,
|
|
72
|
+
FailoverEmbeddingService,
|
|
73
|
+
FailoverImageService,
|
|
74
|
+
FailoverRerankService,
|
|
75
|
+
)
|
|
76
|
+
from kit.ai.pool import AIServicePool
|
|
77
|
+
from kit.ai.profiles import ProfileStatus, get_model_registry
|
|
78
|
+
from kit.ai.rerank_service import RerankHit
|
|
79
|
+
from kit.config.log_config import get_default_logger
|
|
80
|
+
|
|
81
|
+
T = TypeVar("T")
|
|
82
|
+
|
|
83
|
+
logger = get_default_logger(__name__)
|
|
84
|
+
|
|
85
|
+
# 后台常驻服务类型键
|
|
86
|
+
_CHAT = "chat"
|
|
87
|
+
_EMBEDDING = "embedding"
|
|
88
|
+
_RERANK = "rerank"
|
|
89
|
+
_IMAGE = "image"
|
|
90
|
+
|
|
91
|
+
# 支持的模型类型集合(use/available 校验与错误提示用)
|
|
92
|
+
_SERVICE_TYPES: Final[frozenset[str]] = frozenset(
|
|
93
|
+
{_CHAT, _EMBEDDING, _RERANK, _IMAGE}
|
|
94
|
+
)
|
|
95
|
+
|
|
96
|
+
# 常驻链式服务联合类型(按模型类型缓存于后台线程)
|
|
97
|
+
FailoverService = (
|
|
98
|
+
FailoverChatService
|
|
99
|
+
| FailoverEmbeddingService
|
|
100
|
+
| FailoverRerankService
|
|
101
|
+
| FailoverImageService
|
|
102
|
+
)
|
|
103
|
+
|
|
104
|
+
# 同步门面单次调用的默认等待上限(秒),兜底防进程退出/网络挂起;
|
|
105
|
+
# 各类型实际等待上限由 _facade_timeout 按档案 timeout 自适应(+ 余量)。
|
|
106
|
+
_RUN_TIMEOUT = 120.0
|
|
107
|
+
|
|
108
|
+
# 门面等待上限相对模型侧 timeout 的余量(秒)
|
|
109
|
+
_RUN_TIMEOUT_MARGIN: Final[float] = 10.0
|
|
110
|
+
|
|
111
|
+
# ``chat_stream`` 流结束哨兵(放入线程安全队列)
|
|
112
|
+
_SENTINEL = object()
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
class _BackgroundLoop:
|
|
116
|
+
"""共享后台事件循环:单线程常驻,从任意线程安全投递协程。
|
|
117
|
+
|
|
118
|
+
Python 3.13 实现:``asyncio.Runner`` 持有并驱动一个专属事件循环,
|
|
119
|
+
``run_coroutine_threadsafe`` 负责线程间投递与结果同步。
|
|
120
|
+
"""
|
|
121
|
+
|
|
122
|
+
def __init__(self) -> None:
|
|
123
|
+
self._lock = threading.Lock()
|
|
124
|
+
self._started = threading.Event()
|
|
125
|
+
self._thread: threading.Thread | None = None
|
|
126
|
+
self._runner: asyncio.Runner | None = None
|
|
127
|
+
self._loop: asyncio.AbstractEventLoop | None = None
|
|
128
|
+
# 后台 loop 内的 keep_alive 停止信号(线程间经 call_soon_threadsafe 操作)
|
|
129
|
+
self._keep_alive_event: asyncio.Event | None = None
|
|
130
|
+
# 仅在后台线程内读写,主线程只经 shutdown() 间接操作
|
|
131
|
+
self._services: dict[str, FailoverService] = {}
|
|
132
|
+
# AI 候选池(惰性创建;业务线程 AI.use/AI.available 与 loop 线程共享)
|
|
133
|
+
self._pool: AIServicePool | None = None
|
|
134
|
+
# 预热探测任务(进程内仅投递一次,持引用防 GC)
|
|
135
|
+
self._probe_task: asyncio.Task[None] | None = None
|
|
136
|
+
|
|
137
|
+
# ---------- 生命周期 ----------
|
|
138
|
+
|
|
139
|
+
def _ensure_started(self) -> asyncio.AbstractEventLoop:
|
|
140
|
+
loop = self._loop
|
|
141
|
+
if loop is not None and not loop.is_closed():
|
|
142
|
+
return loop
|
|
143
|
+
with self._lock:
|
|
144
|
+
loop = self._loop
|
|
145
|
+
if loop is not None and not loop.is_closed():
|
|
146
|
+
return loop
|
|
147
|
+
self._started = threading.Event()
|
|
148
|
+
self._thread = threading.Thread(
|
|
149
|
+
target=self._worker, name="kit-ai-runner", daemon=True
|
|
150
|
+
)
|
|
151
|
+
self._thread.start()
|
|
152
|
+
# 在 lock 外等待,避免与 worker 内获取 lock 设置 _loop 相互死锁
|
|
153
|
+
self._started.wait(timeout=5)
|
|
154
|
+
loop = self._loop
|
|
155
|
+
if loop is None:
|
|
156
|
+
raise RuntimeError("kit/ai 后台事件循环启动失败")
|
|
157
|
+
return loop
|
|
158
|
+
|
|
159
|
+
def _worker(self) -> None:
|
|
160
|
+
# Runner 需在专属线程内创建(loop 绑定创建线程)
|
|
161
|
+
runner = asyncio.Runner()
|
|
162
|
+
loop = runner.get_loop()
|
|
163
|
+
with self._lock:
|
|
164
|
+
self._runner = runner
|
|
165
|
+
self._loop = loop
|
|
166
|
+
self._keep_alive_event = asyncio.Event()
|
|
167
|
+
self._started.set()
|
|
168
|
+
try:
|
|
169
|
+
# keep_alive 阻塞等待停止信号,让 runner.run 持续处理投递的协程
|
|
170
|
+
runner.run(self._keep_alive())
|
|
171
|
+
except asyncio.CancelledError:
|
|
172
|
+
logger.debug("kit/ai 后台事件循环被取消")
|
|
173
|
+
finally:
|
|
174
|
+
self._close_services()
|
|
175
|
+
with self._lock:
|
|
176
|
+
runner.close()
|
|
177
|
+
self._runner = None
|
|
178
|
+
self._loop = None
|
|
179
|
+
self._keep_alive_event = None
|
|
180
|
+
|
|
181
|
+
async def _keep_alive(self) -> None:
|
|
182
|
+
event = self._keep_alive_event
|
|
183
|
+
if event is not None:
|
|
184
|
+
await event.wait()
|
|
185
|
+
|
|
186
|
+
def _close_services(self) -> None:
|
|
187
|
+
"""在后台线程内关闭候选池中全部常驻服务(关闭底层 AsyncClient)。
|
|
188
|
+
|
|
189
|
+
Failover 链式服务本身不持连接,连接在池内各档案的 service 上,
|
|
190
|
+
统一经 ``pool.close_all()`` 释放。
|
|
191
|
+
"""
|
|
192
|
+
loop = self._loop
|
|
193
|
+
if loop is None or loop.is_closed():
|
|
194
|
+
return
|
|
195
|
+
self._services = {}
|
|
196
|
+
pool = self._pool
|
|
197
|
+
if pool is None:
|
|
198
|
+
return
|
|
199
|
+
try:
|
|
200
|
+
loop.run_until_complete(pool.close_all())
|
|
201
|
+
except Exception as exc: # noqa: BLE001 - 关闭阶段失败不影响整体退出
|
|
202
|
+
logger.warning(f"关闭 AI 候选池失败: {exc}")
|
|
203
|
+
|
|
204
|
+
def shutdown(self) -> None:
|
|
205
|
+
"""停止后台事件循环并释放常驻服务。可在任意线程调用,幂等。"""
|
|
206
|
+
with self._lock:
|
|
207
|
+
loop = self._loop
|
|
208
|
+
thread = self._thread
|
|
209
|
+
event = self._keep_alive_event
|
|
210
|
+
# 触发 keep_alive 结束,让 runner.run 干净返回,避免 pending task 泄漏
|
|
211
|
+
if loop is not None and loop.is_running() and event is not None:
|
|
212
|
+
loop.call_soon_threadsafe(event.set)
|
|
213
|
+
if thread is not None and thread is not threading.current_thread():
|
|
214
|
+
thread.join(timeout=5)
|
|
215
|
+
|
|
216
|
+
# ---------- 候选池与当前项切换 ----------
|
|
217
|
+
|
|
218
|
+
def _get_pool(self) -> AIServicePool:
|
|
219
|
+
"""取(或首次创建)候选池:锁内单例,业务线程与 loop 线程共享。"""
|
|
220
|
+
with self._lock:
|
|
221
|
+
if self._pool is None:
|
|
222
|
+
self._pool = AIServicePool()
|
|
223
|
+
return self._pool
|
|
224
|
+
|
|
225
|
+
def _schedule_probe(self) -> None:
|
|
226
|
+
"""进程内首次取服务时投递候选链预热探测(fire-and-forget,不阻塞)。
|
|
227
|
+
|
|
228
|
+
``probe_all`` 逐类型独立容错且自身不抛异常;持引用防任务被 GC。
|
|
229
|
+
"""
|
|
230
|
+
if self._probe_task is not None:
|
|
231
|
+
return
|
|
232
|
+
pool = self._get_pool()
|
|
233
|
+
self._probe_task = asyncio.create_task(pool.probe_all())
|
|
234
|
+
|
|
235
|
+
def use(self, model_type: str, profile: str) -> None:
|
|
236
|
+
"""切换该类型长期当前档案:指定项成为固定链头并清除冷却。
|
|
237
|
+
|
|
238
|
+
同步方法,可从任意业务线程直接调用(池内部以锁保护共享状态)。
|
|
239
|
+
|
|
240
|
+
Raises:
|
|
241
|
+
AIConfigError: 类型不存在、档案不存在或不可用
|
|
242
|
+
(未配置 / 密钥无效 / 不可达 / 因约束被排除)。
|
|
243
|
+
"""
|
|
244
|
+
self._get_pool().prefer(model_type, profile)
|
|
245
|
+
|
|
246
|
+
def available(self, model_type: str) -> list[ProfileStatus]:
|
|
247
|
+
"""返回该类型全部档案状态快照(含未配置者),链头标记 current。"""
|
|
248
|
+
return self._get_pool().available(model_type)
|
|
249
|
+
|
|
250
|
+
# ---------- 投递 ----------
|
|
251
|
+
|
|
252
|
+
def run(
|
|
253
|
+
self,
|
|
254
|
+
coro_factory: Callable[[], Coroutine[Any, Any, T]],
|
|
255
|
+
*,
|
|
256
|
+
timeout: float = _RUN_TIMEOUT,
|
|
257
|
+
) -> T:
|
|
258
|
+
"""同步执行协程工厂:投递到后台 loop 并阻塞取结果。
|
|
259
|
+
|
|
260
|
+
因使用独立后台 loop,从同步代码、线程内均可安全调用
|
|
261
|
+
(``run_coroutine_threadsafe`` 线程安全)。
|
|
262
|
+
|
|
263
|
+
Args:
|
|
264
|
+
coro_factory: 协程工厂(每次调用产出一个新协程)
|
|
265
|
+
timeout: 同步等待上限(秒);应不低于协程内部模型侧超时
|
|
266
|
+
"""
|
|
267
|
+
loop = self._ensure_started()
|
|
268
|
+
if threading.current_thread() is self._thread:
|
|
269
|
+
# 避免在后台线程内二次同步等待造成死锁
|
|
270
|
+
raise RuntimeError("kit/ai 同步门面不支持在后台线程内再次同步调用")
|
|
271
|
+
coro = coro_factory()
|
|
272
|
+
future: concurrent.futures.Future[T] = asyncio.run_coroutine_threadsafe(
|
|
273
|
+
coro, loop
|
|
274
|
+
)
|
|
275
|
+
try:
|
|
276
|
+
return future.result(timeout=timeout)
|
|
277
|
+
except concurrent.futures.TimeoutError:
|
|
278
|
+
# 后台 loop 内的协程仍在运行,取消它避免泄漏
|
|
279
|
+
future.cancel()
|
|
280
|
+
raise TimeoutError(f"kit/ai 同步调用超过 {timeout:g}s 未返回") from None
|
|
281
|
+
|
|
282
|
+
def submit_stream(
|
|
283
|
+
self,
|
|
284
|
+
messages: list[dict[str, str]],
|
|
285
|
+
q: queue.Queue[object],
|
|
286
|
+
) -> concurrent.futures.Future[None]:
|
|
287
|
+
"""投递流式收集任务:后台拉取链式 ``stream`` 帧写入线程安全队列。
|
|
288
|
+
|
|
289
|
+
返回的 future 可取消,用于同步侧中断时停止后台拉取。
|
|
290
|
+
"""
|
|
291
|
+
loop = self._ensure_started()
|
|
292
|
+
|
|
293
|
+
async def _producer() -> None:
|
|
294
|
+
service: FailoverChatService = cast(
|
|
295
|
+
FailoverChatService, await self._get_service(_CHAT)
|
|
296
|
+
)
|
|
297
|
+
try:
|
|
298
|
+
async for delta in service.stream(messages):
|
|
299
|
+
q.put(delta)
|
|
300
|
+
except BaseException as exc: # noqa: BLE001 - 异常经队列回传给同步侧
|
|
301
|
+
q.put(exc)
|
|
302
|
+
finally:
|
|
303
|
+
q.put(_SENTINEL)
|
|
304
|
+
|
|
305
|
+
return asyncio.run_coroutine_threadsafe(_producer(), loop)
|
|
306
|
+
|
|
307
|
+
async def _get_service(self, service_type: str) -> FailoverService:
|
|
308
|
+
"""后台线程内取常驻链式服务,未创建则就地构建并缓存。
|
|
309
|
+
|
|
310
|
+
进程内首次构建任意类型时投递一次候选链预热探测(fire-and-forget)。
|
|
311
|
+
"""
|
|
312
|
+
service = self._services.get(service_type)
|
|
313
|
+
if service is None:
|
|
314
|
+
self._schedule_probe()
|
|
315
|
+
pool = self._get_pool()
|
|
316
|
+
match service_type:
|
|
317
|
+
case "chat":
|
|
318
|
+
service = FailoverChatService(pool)
|
|
319
|
+
case "embedding":
|
|
320
|
+
service = FailoverEmbeddingService(pool)
|
|
321
|
+
case "rerank":
|
|
322
|
+
service = FailoverRerankService(pool)
|
|
323
|
+
case "image":
|
|
324
|
+
service = FailoverImageService(pool)
|
|
325
|
+
case _:
|
|
326
|
+
raise ValueError(
|
|
327
|
+
f"未知模型类型: {service_type},"
|
|
328
|
+
f"支持: {', '.join(sorted(_SERVICE_TYPES))}"
|
|
329
|
+
)
|
|
330
|
+
self._services[service_type] = service
|
|
331
|
+
logger.debug(f"已构建并常驻 AI 链式服务:{service_type}")
|
|
332
|
+
return service
|
|
333
|
+
|
|
334
|
+
|
|
335
|
+
_background = _BackgroundLoop()
|
|
336
|
+
|
|
337
|
+
|
|
338
|
+
def run_sync[T](
|
|
339
|
+
coro_factory: Callable[[], Coroutine[Any, Any, T]],
|
|
340
|
+
*,
|
|
341
|
+
timeout: float = _RUN_TIMEOUT,
|
|
342
|
+
) -> T:
|
|
343
|
+
"""在任意上下文同步执行协程工厂(兼容旧接口,等价于后台投递)。"""
|
|
344
|
+
return _background.run(coro_factory, timeout=timeout)
|
|
345
|
+
|
|
346
|
+
|
|
347
|
+
def _facade_timeout(model_type: str) -> float:
|
|
348
|
+
"""门面同步等待上限:该类型档案 timeout 最大值 + 余量(秒)。
|
|
349
|
+
|
|
350
|
+
读静态档案(``type_profiles``,无需密钥),随 ``ai_models.json``
|
|
351
|
+
的超时配置自适应;类型未配置时回退默认 ``_RUN_TIMEOUT``。
|
|
352
|
+
"""
|
|
353
|
+
try:
|
|
354
|
+
profiles = get_model_registry().type_profiles(model_type)
|
|
355
|
+
except AIConfigError:
|
|
356
|
+
return _RUN_TIMEOUT
|
|
357
|
+
if not profiles:
|
|
358
|
+
return _RUN_TIMEOUT
|
|
359
|
+
model_timeout = max(config.timeout for config in profiles.values())
|
|
360
|
+
return float(model_timeout) + _RUN_TIMEOUT_MARGIN
|
|
361
|
+
|
|
362
|
+
|
|
363
|
+
async def _chat_ask(messages: list[dict[str, str]]) -> str:
|
|
364
|
+
service: FailoverChatService = cast(
|
|
365
|
+
FailoverChatService, await _background._get_service(_CHAT) # noqa: SLF001
|
|
366
|
+
)
|
|
367
|
+
return await service.ask(messages)
|
|
368
|
+
|
|
369
|
+
|
|
370
|
+
async def _chat_complete(prompt: str) -> str:
|
|
371
|
+
service: FailoverChatService = cast(
|
|
372
|
+
FailoverChatService, await _background._get_service(_CHAT) # noqa: SLF001
|
|
373
|
+
)
|
|
374
|
+
return await service.complete(prompt)
|
|
375
|
+
|
|
376
|
+
|
|
377
|
+
async def _embed_texts(texts: list[str]) -> list[list[float]]:
|
|
378
|
+
if not texts:
|
|
379
|
+
return []
|
|
380
|
+
service: FailoverEmbeddingService = cast(
|
|
381
|
+
FailoverEmbeddingService,
|
|
382
|
+
await _background._get_service(_EMBEDDING), # noqa: SLF001
|
|
383
|
+
)
|
|
384
|
+
return await service.embed_texts(texts)
|
|
385
|
+
|
|
386
|
+
|
|
387
|
+
async def _rerank(
|
|
388
|
+
query: str, documents: list[str], top_n: int | None
|
|
389
|
+
) -> list[tuple[int, float]]:
|
|
390
|
+
if not documents:
|
|
391
|
+
return []
|
|
392
|
+
service: FailoverRerankService = cast(
|
|
393
|
+
FailoverRerankService,
|
|
394
|
+
await _background._get_service(_RERANK), # noqa: SLF001
|
|
395
|
+
)
|
|
396
|
+
hits: list[RerankHit] = await service.rerank(query, documents, top_n=top_n)
|
|
397
|
+
return [(hit.index, hit.score) for hit in hits]
|
|
398
|
+
|
|
399
|
+
|
|
400
|
+
async def _generate_image(prompt: str, *, size: str | None, n: int) -> str:
|
|
401
|
+
service: FailoverImageService = cast(
|
|
402
|
+
FailoverImageService,
|
|
403
|
+
await _background._get_service(_IMAGE), # noqa: SLF001
|
|
404
|
+
)
|
|
405
|
+
return await service.generate(prompt, size=size, n=n)
|
|
406
|
+
|
|
407
|
+
|
|
408
|
+
class AI:
|
|
409
|
+
"""kit/ai 同步门面:在同步上下文调用异步 AI 服务(候选链故障切换)。"""
|
|
410
|
+
|
|
411
|
+
@staticmethod
|
|
412
|
+
def chat_complete(prompt: str) -> str:
|
|
413
|
+
"""单轮对话(仅 user 消息),返回模型回答文本。"""
|
|
414
|
+
return _background.run(
|
|
415
|
+
lambda: _chat_complete(prompt), timeout=_facade_timeout(_CHAT)
|
|
416
|
+
)
|
|
417
|
+
|
|
418
|
+
@staticmethod
|
|
419
|
+
def chat_ask(messages: list[dict[str, str]]) -> str:
|
|
420
|
+
"""多轮对话(自组装 messages),返回模型回答文本。"""
|
|
421
|
+
return _background.run(
|
|
422
|
+
lambda: _chat_ask(messages), timeout=_facade_timeout(_CHAT)
|
|
423
|
+
)
|
|
424
|
+
|
|
425
|
+
@staticmethod
|
|
426
|
+
def embed_text(text: str) -> list[float]:
|
|
427
|
+
"""对单个文本向量化,返回向量。"""
|
|
428
|
+
return _background.run(
|
|
429
|
+
lambda: _embed_texts([text]), timeout=_facade_timeout(_EMBEDDING)
|
|
430
|
+
)[0]
|
|
431
|
+
|
|
432
|
+
@staticmethod
|
|
433
|
+
def embed_texts(texts: list[str]) -> list[list[float]]:
|
|
434
|
+
"""对文本列表批量向量化,返回向量列表。"""
|
|
435
|
+
return _background.run(
|
|
436
|
+
lambda: _embed_texts(texts), timeout=_facade_timeout(_EMBEDDING)
|
|
437
|
+
)
|
|
438
|
+
|
|
439
|
+
@staticmethod
|
|
440
|
+
def rerank(
|
|
441
|
+
query: str, documents: list[str], top_n: int | None = None
|
|
442
|
+
) -> list[tuple[int, float]]:
|
|
443
|
+
"""对候选文档重排序,返回 [(原索引, 相关分)] 降序列表。"""
|
|
444
|
+
return _background.run(
|
|
445
|
+
lambda: _rerank(query, documents, top_n), timeout=_facade_timeout(_RERANK)
|
|
446
|
+
)
|
|
447
|
+
|
|
448
|
+
@staticmethod
|
|
449
|
+
def generate_image(prompt: str, *, size: str | None = None, n: int = 1) -> str:
|
|
450
|
+
"""文生图,返回生成的图片 URL。"""
|
|
451
|
+
return _background.run(
|
|
452
|
+
lambda: _generate_image(prompt, size=size, n=n),
|
|
453
|
+
timeout=_facade_timeout(_IMAGE),
|
|
454
|
+
)
|
|
455
|
+
|
|
456
|
+
@staticmethod
|
|
457
|
+
def use(model_type: str, profile: str) -> None:
|
|
458
|
+
"""切换某类型长期当前模型档案(进程内生效,重启回 env/json 序)。
|
|
459
|
+
|
|
460
|
+
指定档案成为固定链头并清除其冷却;再次 ``use`` 才会更换。
|
|
461
|
+
|
|
462
|
+
Args:
|
|
463
|
+
model_type: 模型类型(chat / embedding / rerank / image)
|
|
464
|
+
profile: 目标档案名(``ai_models.json`` 中该类型 profiles 的键)
|
|
465
|
+
|
|
466
|
+
Raises:
|
|
467
|
+
AIConfigError: 类型不存在、档案不存在或不可用
|
|
468
|
+
(未配置 / 密钥无效 / 不可达 / 因约束被排除)。
|
|
469
|
+
"""
|
|
470
|
+
_background.use(model_type, profile)
|
|
471
|
+
|
|
472
|
+
@staticmethod
|
|
473
|
+
def available(model_type: str) -> list[ProfileStatus]:
|
|
474
|
+
"""返回某类型全部档案状态快照(含未配置与被剔除者),链头标 current。
|
|
475
|
+
|
|
476
|
+
Args:
|
|
477
|
+
model_type: 模型类型(chat / embedding / rerank / image)
|
|
478
|
+
|
|
479
|
+
Returns:
|
|
480
|
+
``ProfileStatus`` 列表(name/model/state/current/last_error/
|
|
481
|
+
last_ok_at),供业务展示与切换决策。
|
|
482
|
+
|
|
483
|
+
Raises:
|
|
484
|
+
AIConfigError: 类型不存在或注册表未加载。
|
|
485
|
+
"""
|
|
486
|
+
return _background.available(model_type)
|
|
487
|
+
|
|
488
|
+
@staticmethod
|
|
489
|
+
def chat_stream(
|
|
490
|
+
messages: list[dict[str, str]],
|
|
491
|
+
) -> Iterator[str]:
|
|
492
|
+
"""多轮对话(流式),返回逐 token 迭代器(同步)。
|
|
493
|
+
|
|
494
|
+
在同步线程内逐帧消费,适合流式输出;内部在后台 loop 持续拉取
|
|
495
|
+
并放入线程安全队列。
|
|
496
|
+
|
|
497
|
+
推荐配合 ``with`` 使用,保证中途 ``break`` 也会及时取消后台任务::
|
|
498
|
+
|
|
499
|
+
with AI.chat_stream(messages) as stream:
|
|
500
|
+
for token in stream:
|
|
501
|
+
yield token
|
|
502
|
+
if done:
|
|
503
|
+
break
|
|
504
|
+
"""
|
|
505
|
+
return _StreamIterator(messages)
|
|
506
|
+
|
|
507
|
+
|
|
508
|
+
class _StreamIterator:
|
|
509
|
+
"""``AI.chat_stream`` 的同步流式迭代器。
|
|
510
|
+
|
|
511
|
+
通过 ``queue.Queue`` 桥接后台 loop:一个后台任务把链式 ``stream``
|
|
512
|
+
产出的帧逐条放入队列,同步侧阻塞 ``get`` 消费;停止信号由哨兵标记。
|
|
513
|
+
|
|
514
|
+
支持 ``with`` 上下文(推荐):退出时主动取消后台任务;即便调用方
|
|
515
|
+
直接 ``for`` 中途 ``break`` 或迭代器被 GC,``close``/``__del__`` 也会
|
|
516
|
+
兜底取消后台 producer,避免任务残留。
|
|
517
|
+
"""
|
|
518
|
+
|
|
519
|
+
def __init__(self, messages: list[dict[str, str]]) -> None:
|
|
520
|
+
self._messages: list[dict[str, str]] = messages
|
|
521
|
+
self._q: queue.Queue[object] = queue.Queue()
|
|
522
|
+
self._started = False
|
|
523
|
+
self._closed = False
|
|
524
|
+
self._task: concurrent.futures.Future[None] | None = None
|
|
525
|
+
|
|
526
|
+
def __iter__(self) -> Iterator[str]:
|
|
527
|
+
return self
|
|
528
|
+
|
|
529
|
+
def _ensure_started(self) -> None:
|
|
530
|
+
if not self._started:
|
|
531
|
+
self._started = True
|
|
532
|
+
self._task = _background.submit_stream(self._messages, self._q)
|
|
533
|
+
|
|
534
|
+
def _cancel_task(self) -> None:
|
|
535
|
+
task = self._task
|
|
536
|
+
if task is not None:
|
|
537
|
+
task.cancel()
|
|
538
|
+
self._task = None
|
|
539
|
+
|
|
540
|
+
def close(self) -> None:
|
|
541
|
+
"""停止后台拉取并释放资源。幂等,可多次调用。"""
|
|
542
|
+
if self._closed:
|
|
543
|
+
return
|
|
544
|
+
self._closed = True
|
|
545
|
+
self._cancel_task()
|
|
546
|
+
|
|
547
|
+
def __del__(self) -> None:
|
|
548
|
+
# 兜底清理:for 中途 break / 迭代器被 GC 时取消后台 producer。
|
|
549
|
+
# 模块卸载期 _task 是独立 future 引用,cancel 对已 shutdown 的 loop 安全。
|
|
550
|
+
try:
|
|
551
|
+
self.close()
|
|
552
|
+
except Exception: # noqa: BLE001 - GC 阶段不容许抛异常
|
|
553
|
+
pass
|
|
554
|
+
|
|
555
|
+
def __enter__(self) -> _StreamIterator:
|
|
556
|
+
return self
|
|
557
|
+
|
|
558
|
+
def __exit__(self, *exc_info: object) -> None:
|
|
559
|
+
self.close()
|
|
560
|
+
|
|
561
|
+
def __next__(self) -> str:
|
|
562
|
+
if self._closed:
|
|
563
|
+
raise StopIteration
|
|
564
|
+
self._ensure_started()
|
|
565
|
+
item = self._q.get()
|
|
566
|
+
if item is _SENTINEL:
|
|
567
|
+
self.close()
|
|
568
|
+
raise StopIteration
|
|
569
|
+
if isinstance(item, BaseException):
|
|
570
|
+
self.close()
|
|
571
|
+
raise item
|
|
572
|
+
return str(item)
|
|
573
|
+
|
|
574
|
+
|
|
575
|
+
def shutdown() -> None:
|
|
576
|
+
"""停止后台事件循环并释放常驻 AI 服务(应用退出前调用,幂等)。"""
|
|
577
|
+
_background.shutdown()
|
|
578
|
+
|
|
579
|
+
|
|
580
|
+
__all__ = ["AI", "run_sync", "shutdown"]
|
kit/ai/types.py
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
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: sdpy
|
|
17
|
+
Author: Clark Chang
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
"""kit/ai 公共类型别名(PEP 695 type 语句,递归 JSON 类型)"""
|
|
21
|
+
|
|
22
|
+
# 递归 JSON 值类型:标量 | 数组 | 对象
|
|
23
|
+
type JSONScalar = str | int | float | bool | None
|
|
24
|
+
type JSONValue = JSONScalar | list[JSONValue] | dict[str, JSONValue]
|
|
25
|
+
# JSON 对象(请求体 / 响应体)
|
|
26
|
+
type JSONObject = dict[str, JSONValue]
|
|
27
|
+
|
|
28
|
+
__all__ = ["JSONObject", "JSONScalar", "JSONValue"]
|