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/failover.py
ADDED
|
@@ -0,0 +1,244 @@
|
|
|
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
|
+
"""候选链故障切换调用层
|
|
22
|
+
|
|
23
|
+
将单档案 Service 的公开方法包装为链式调用:按池内候选顺序尝试,任一档案
|
|
24
|
+
抛出 ``AIError`` 即冷却该档案并尝试下一个;全部失败抛出携带已尝试档案
|
|
25
|
+
列表与最后失败原因的 ``AIRequestError``。
|
|
26
|
+
|
|
27
|
+
流式对话仅在「首帧到达前」允许切换档案;首帧已产出后失败直接抛出
|
|
28
|
+
(已输出 token 无法回滚)。空流(首帧即结束)视为正常完成。
|
|
29
|
+
|
|
30
|
+
各 Failover 服务公开方法与对应单档案 Service 保持同形,业务侧无感替换。
|
|
31
|
+
"""
|
|
32
|
+
|
|
33
|
+
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable
|
|
34
|
+
from typing import TypeVar, cast
|
|
35
|
+
|
|
36
|
+
from kit.ai.chat_service import ChatService
|
|
37
|
+
from kit.ai.embedding_service import EmbeddingService
|
|
38
|
+
from kit.ai.exceptions import AIError, AIRequestError
|
|
39
|
+
from kit.ai.image_service import ImageService
|
|
40
|
+
from kit.ai.model_builder import Service
|
|
41
|
+
from kit.ai.pool import AIServicePool
|
|
42
|
+
from kit.ai.rerank_service import RerankHit, RerankService
|
|
43
|
+
|
|
44
|
+
__all__ = [
|
|
45
|
+
"FailoverChatService",
|
|
46
|
+
"FailoverEmbeddingService",
|
|
47
|
+
"FailoverImageService",
|
|
48
|
+
"FailoverRerankService",
|
|
49
|
+
]
|
|
50
|
+
|
|
51
|
+
T = TypeVar("T")
|
|
52
|
+
|
|
53
|
+
# 批量向量化默认并发(与 EmbeddingService 保持一致)
|
|
54
|
+
_EMBED_CONCURRENCY = 4
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
class _FailoverBase:
|
|
58
|
+
"""链式调用基类:候选迭代 + 失败冷却 + 聚合异常。"""
|
|
59
|
+
|
|
60
|
+
def __init__(self, pool: AIServicePool, model_type: str) -> None:
|
|
61
|
+
self._pool: AIServicePool = pool
|
|
62
|
+
self._model_type: str = model_type
|
|
63
|
+
|
|
64
|
+
async def _attempt(
|
|
65
|
+
self,
|
|
66
|
+
label: str,
|
|
67
|
+
run: Callable[[Service], Awaitable[T]],
|
|
68
|
+
) -> T:
|
|
69
|
+
"""按候选顺序执行 ``run``;单档案失败冷却后换下一个,全失败聚合抛出。"""
|
|
70
|
+
tried: list[str] = []
|
|
71
|
+
last_exc: AIError | None = None
|
|
72
|
+
for name, service in self._pool.iter_services(self._model_type):
|
|
73
|
+
tried.append(name)
|
|
74
|
+
try:
|
|
75
|
+
result = await run(service)
|
|
76
|
+
except AIError as exc:
|
|
77
|
+
last_exc = exc
|
|
78
|
+
self._pool.mark_failure(self._model_type, name, exc)
|
|
79
|
+
continue
|
|
80
|
+
self._pool.mark_success(self._model_type, name)
|
|
81
|
+
return result
|
|
82
|
+
raise AIRequestError(
|
|
83
|
+
f"{label}失败:已尝试 {len(tried)} 个档案"
|
|
84
|
+
f"({', '.join(tried) or '无可用候选'}),最后错误: {last_exc}"
|
|
85
|
+
) from last_exc
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
class FailoverChatService(_FailoverBase):
|
|
89
|
+
"""对话服务链式调用(ask / complete / stream)。"""
|
|
90
|
+
|
|
91
|
+
def __init__(self, pool: AIServicePool) -> None:
|
|
92
|
+
super().__init__(pool, "chat")
|
|
93
|
+
|
|
94
|
+
async def ask(
|
|
95
|
+
self,
|
|
96
|
+
messages: list[dict[str, str]],
|
|
97
|
+
*,
|
|
98
|
+
temperature: float | None = None,
|
|
99
|
+
max_tokens: int | None = None,
|
|
100
|
+
) -> str:
|
|
101
|
+
"""多轮对话,返回模型回答文本(链式故障切换)。"""
|
|
102
|
+
|
|
103
|
+
async def _run(service: Service) -> str:
|
|
104
|
+
return await cast(ChatService, service).ask(
|
|
105
|
+
messages, temperature=temperature, max_tokens=max_tokens
|
|
106
|
+
)
|
|
107
|
+
|
|
108
|
+
return await self._attempt("对话", _run)
|
|
109
|
+
|
|
110
|
+
async def complete(
|
|
111
|
+
self,
|
|
112
|
+
prompt: str,
|
|
113
|
+
*,
|
|
114
|
+
temperature: float | None = None,
|
|
115
|
+
max_tokens: int | None = None,
|
|
116
|
+
) -> str:
|
|
117
|
+
"""单轮对话,返回模型回答文本(链式故障切换)。"""
|
|
118
|
+
|
|
119
|
+
async def _run(service: Service) -> str:
|
|
120
|
+
return await cast(ChatService, service).complete(
|
|
121
|
+
prompt, temperature=temperature, max_tokens=max_tokens
|
|
122
|
+
)
|
|
123
|
+
|
|
124
|
+
return await self._attempt("单轮对话", _run)
|
|
125
|
+
|
|
126
|
+
async def stream(
|
|
127
|
+
self,
|
|
128
|
+
messages: list[dict[str, str]],
|
|
129
|
+
*,
|
|
130
|
+
temperature: float | None = None,
|
|
131
|
+
max_tokens: int | None = None,
|
|
132
|
+
) -> AsyncIterator[str]:
|
|
133
|
+
"""流式对话:首帧前失败自动换档案重试;首帧后失败直接抛出。"""
|
|
134
|
+
tried: list[str] = []
|
|
135
|
+
last_exc: AIError | None = None
|
|
136
|
+
for name, service in self._pool.iter_services(self._model_type):
|
|
137
|
+
tried.append(name)
|
|
138
|
+
agen = cast(
|
|
139
|
+
AsyncGenerator[str, None],
|
|
140
|
+
cast(ChatService, service).stream(
|
|
141
|
+
messages, temperature=temperature, max_tokens=max_tokens
|
|
142
|
+
),
|
|
143
|
+
)
|
|
144
|
+
try:
|
|
145
|
+
first = await agen.__anext__()
|
|
146
|
+
except StopAsyncIteration:
|
|
147
|
+
# 空流(首帧即结束)视为正常完成
|
|
148
|
+
await agen.aclose()
|
|
149
|
+
return
|
|
150
|
+
except AIError as exc:
|
|
151
|
+
last_exc = exc
|
|
152
|
+
self._pool.mark_failure(self._model_type, name, exc)
|
|
153
|
+
await agen.aclose()
|
|
154
|
+
continue
|
|
155
|
+
try:
|
|
156
|
+
self._pool.mark_success(self._model_type, name)
|
|
157
|
+
yield first
|
|
158
|
+
async for delta in agen:
|
|
159
|
+
yield delta
|
|
160
|
+
finally:
|
|
161
|
+
await agen.aclose()
|
|
162
|
+
return
|
|
163
|
+
raise AIRequestError(
|
|
164
|
+
f"流式对话失败:已尝试 {len(tried)} 个档案"
|
|
165
|
+
f"({', '.join(tried) or '无可用候选'}),最后错误: {last_exc}"
|
|
166
|
+
) from last_exc
|
|
167
|
+
|
|
168
|
+
|
|
169
|
+
class FailoverEmbeddingService(_FailoverBase):
|
|
170
|
+
"""向量化服务链式调用(embed_text / embed_texts)。"""
|
|
171
|
+
|
|
172
|
+
def __init__(self, pool: AIServicePool) -> None:
|
|
173
|
+
super().__init__(pool, "embedding")
|
|
174
|
+
|
|
175
|
+
async def embed_text(self, text: str) -> list[float]:
|
|
176
|
+
"""单文本向量化(链式故障切换)。"""
|
|
177
|
+
vectors = await self.embed_texts([text])
|
|
178
|
+
return vectors[0]
|
|
179
|
+
|
|
180
|
+
async def embed_texts(
|
|
181
|
+
self,
|
|
182
|
+
texts: list[str],
|
|
183
|
+
*,
|
|
184
|
+
batch_size: int | None = None,
|
|
185
|
+
concurrency: int = _EMBED_CONCURRENCY,
|
|
186
|
+
) -> list[list[float]]:
|
|
187
|
+
"""批量向量化(链式故障切换;整批失败才切换档案)。"""
|
|
188
|
+
if not texts:
|
|
189
|
+
return []
|
|
190
|
+
|
|
191
|
+
async def _run(service: Service) -> list[list[float]]:
|
|
192
|
+
return await cast(EmbeddingService, service).embed_texts(
|
|
193
|
+
texts, batch_size=batch_size, concurrency=concurrency
|
|
194
|
+
)
|
|
195
|
+
|
|
196
|
+
return await self._attempt("向量化", _run)
|
|
197
|
+
|
|
198
|
+
|
|
199
|
+
class FailoverRerankService(_FailoverBase):
|
|
200
|
+
"""重排序服务链式调用(rerank)。"""
|
|
201
|
+
|
|
202
|
+
def __init__(self, pool: AIServicePool) -> None:
|
|
203
|
+
super().__init__(pool, "rerank")
|
|
204
|
+
|
|
205
|
+
async def rerank(
|
|
206
|
+
self,
|
|
207
|
+
query: str,
|
|
208
|
+
documents: list[str],
|
|
209
|
+
*,
|
|
210
|
+
top_n: int | None = None,
|
|
211
|
+
) -> list[RerankHit]:
|
|
212
|
+
"""候选文档重排序(链式故障切换)。"""
|
|
213
|
+
if not documents:
|
|
214
|
+
return []
|
|
215
|
+
|
|
216
|
+
async def _run(service: Service) -> list[RerankHit]:
|
|
217
|
+
return await cast(RerankService, service).rerank(
|
|
218
|
+
query, documents, top_n=top_n
|
|
219
|
+
)
|
|
220
|
+
|
|
221
|
+
return await self._attempt("重排序", _run)
|
|
222
|
+
|
|
223
|
+
|
|
224
|
+
class FailoverImageService(_FailoverBase):
|
|
225
|
+
"""图像生成服务链式调用(generate)。"""
|
|
226
|
+
|
|
227
|
+
def __init__(self, pool: AIServicePool) -> None:
|
|
228
|
+
super().__init__(pool, "image")
|
|
229
|
+
|
|
230
|
+
async def generate(
|
|
231
|
+
self,
|
|
232
|
+
prompt: str,
|
|
233
|
+
*,
|
|
234
|
+
size: str | None = None,
|
|
235
|
+
n: int = 1,
|
|
236
|
+
) -> str:
|
|
237
|
+
"""文生图,返回图片 URL(链式故障切换)。"""
|
|
238
|
+
|
|
239
|
+
async def _run(service: Service) -> str:
|
|
240
|
+
return await cast(ImageService, service).generate(
|
|
241
|
+
prompt, size=size, n=n
|
|
242
|
+
)
|
|
243
|
+
|
|
244
|
+
return await self._attempt("图像生成", _run)
|
kit/ai/image_service.py
ADDED
|
@@ -0,0 +1,99 @@
|
|
|
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 dataclasses import dataclass
|
|
21
|
+
from typing import override
|
|
22
|
+
|
|
23
|
+
from kit.ai.base_service import BaseAIService, safe_get
|
|
24
|
+
from kit.ai.exceptions import AIResponseError
|
|
25
|
+
from kit.ai.types import JSONObject
|
|
26
|
+
|
|
27
|
+
__all__ = ["ImageRequest", "ImageService"]
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
@dataclass(frozen=True)
|
|
31
|
+
class ImageRequest:
|
|
32
|
+
"""文生图请求模型。
|
|
33
|
+
|
|
34
|
+
``return_base64`` 对应 Agnes 官方参数 ``return_base64``,为 ``True``
|
|
35
|
+
时服务端以 ``data[0].b64_json`` 返回 Base64 数据而非 URL。
|
|
36
|
+
``ratio`` 与档位式 ``size`` 配合的宽高比,官方支持
|
|
37
|
+
``1:1/3:4/4:3/16:9/9:16/2:3/3:2/21:9``,默认 ``1:1``。
|
|
38
|
+
"""
|
|
39
|
+
|
|
40
|
+
prompt: str
|
|
41
|
+
size: str | None = None
|
|
42
|
+
n: int = 1
|
|
43
|
+
return_base64: bool = False
|
|
44
|
+
ratio: str | None = None
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
class ImageService(BaseAIService[ImageRequest, str]):
|
|
48
|
+
"""图像生成服务(全异步)。
|
|
49
|
+
|
|
50
|
+
``n > 1`` 时仅返回第一张图片,调用方如需多图请自行扩展。
|
|
51
|
+
返回 ``str`` 语义:默认返回图片 URL;``return_base64=True`` 时
|
|
52
|
+
返回 Base64 数据,调用方按返回值是否以 ``http`` 开头区分形式。
|
|
53
|
+
"""
|
|
54
|
+
|
|
55
|
+
@override
|
|
56
|
+
def _build_payload(self, request: ImageRequest) -> JSONObject:
|
|
57
|
+
payload: JSONObject = {
|
|
58
|
+
"model": self._profile.model,
|
|
59
|
+
"prompt": request.prompt,
|
|
60
|
+
"n": request.n,
|
|
61
|
+
}
|
|
62
|
+
if request.size is not None:
|
|
63
|
+
payload["size"] = request.size
|
|
64
|
+
if request.return_base64:
|
|
65
|
+
payload["return_base64"] = True
|
|
66
|
+
if request.ratio is not None:
|
|
67
|
+
payload["ratio"] = request.ratio
|
|
68
|
+
return payload
|
|
69
|
+
|
|
70
|
+
@override
|
|
71
|
+
def _parse_response(self, data: JSONObject) -> str:
|
|
72
|
+
# 优先 URL 形式(data[0].url);Base64 形式下 url 为 null 或整个键缺失
|
|
73
|
+
# (如 sensenova 响应仅含 b64_json),回退取 data[0].b64_json,
|
|
74
|
+
# 两者均无效才判定响应异常。
|
|
75
|
+
first = safe_get(data, "data", 0)
|
|
76
|
+
if isinstance(first, dict):
|
|
77
|
+
url = first.get("url")
|
|
78
|
+
if isinstance(url, str):
|
|
79
|
+
return url
|
|
80
|
+
b64 = first.get("b64_json")
|
|
81
|
+
if isinstance(b64, str):
|
|
82
|
+
return b64
|
|
83
|
+
raise AIResponseError(
|
|
84
|
+
f"图像生成响应缺少 url/b64_json: {str(data)[:200]}"
|
|
85
|
+
)
|
|
86
|
+
|
|
87
|
+
async def generate(
|
|
88
|
+
self,
|
|
89
|
+
prompt: str,
|
|
90
|
+
*,
|
|
91
|
+
size: str | None = None,
|
|
92
|
+
n: int = 1,
|
|
93
|
+
return_base64: bool = False,
|
|
94
|
+
ratio: str | None = None,
|
|
95
|
+
) -> str:
|
|
96
|
+
request = ImageRequest(
|
|
97
|
+
prompt, size=size, n=n, return_base64=return_base64, ratio=ratio
|
|
98
|
+
)
|
|
99
|
+
return await self._call(request)
|
kit/ai/model_builder.py
ADDED
|
@@ -0,0 +1,119 @@
|
|
|
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 kit.ai.chat_service import ChatService
|
|
21
|
+
from kit.ai.embedding_service import EmbeddingService
|
|
22
|
+
from kit.ai.image_service import ImageService
|
|
23
|
+
from kit.ai.profiles import ModelProfile, get_model_registry, with_api_key
|
|
24
|
+
from kit.ai.rerank_service import RerankService
|
|
25
|
+
|
|
26
|
+
__all__ = ["ModelBuilder"]
|
|
27
|
+
|
|
28
|
+
# 支持的模型类型
|
|
29
|
+
_MODEL_TYPES: tuple[str, ...] = ("chat", "embedding", "rerank", "image")
|
|
30
|
+
|
|
31
|
+
Service = ChatService | EmbeddingService | RerankService | ImageService
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
class ModelBuilder:
|
|
35
|
+
"""根据模型类型构建对应的 AI 服务实例。"""
|
|
36
|
+
|
|
37
|
+
@staticmethod
|
|
38
|
+
def _resolve_profile(
|
|
39
|
+
model_type: str,
|
|
40
|
+
model: str | None = None,
|
|
41
|
+
*,
|
|
42
|
+
api_key: str | None = None,
|
|
43
|
+
) -> ModelProfile:
|
|
44
|
+
registry = get_model_registry()
|
|
45
|
+
profile = registry.get(model_type, model)
|
|
46
|
+
if api_key is not None:
|
|
47
|
+
profile = with_api_key(profile, api_key)
|
|
48
|
+
return profile
|
|
49
|
+
|
|
50
|
+
@staticmethod
|
|
51
|
+
def build(
|
|
52
|
+
model_type: str,
|
|
53
|
+
model: str | None = None,
|
|
54
|
+
*,
|
|
55
|
+
api_key: str | None = None,
|
|
56
|
+
) -> Service:
|
|
57
|
+
"""构建指定的 AI 服务。
|
|
58
|
+
|
|
59
|
+
Args:
|
|
60
|
+
model_type: 模型类型(chat / embedding / rerank / image)
|
|
61
|
+
model: 档案名,None 时回退到该类型默认档案
|
|
62
|
+
api_key: 可选,覆盖档案中的密钥
|
|
63
|
+
|
|
64
|
+
Raises:
|
|
65
|
+
ValueError: 未知模型类型。
|
|
66
|
+
"""
|
|
67
|
+
match model_type:
|
|
68
|
+
case "chat":
|
|
69
|
+
profile = ModelBuilder._resolve_profile(
|
|
70
|
+
model_type, model, api_key=api_key
|
|
71
|
+
)
|
|
72
|
+
return ChatService(profile)
|
|
73
|
+
case "embedding":
|
|
74
|
+
profile = ModelBuilder._resolve_profile(
|
|
75
|
+
model_type, model, api_key=api_key
|
|
76
|
+
)
|
|
77
|
+
return EmbeddingService(profile)
|
|
78
|
+
case "rerank":
|
|
79
|
+
profile = ModelBuilder._resolve_profile(
|
|
80
|
+
model_type, model, api_key=api_key
|
|
81
|
+
)
|
|
82
|
+
return RerankService(profile)
|
|
83
|
+
case "image":
|
|
84
|
+
profile = ModelBuilder._resolve_profile(
|
|
85
|
+
model_type, model, api_key=api_key
|
|
86
|
+
)
|
|
87
|
+
return ImageService(profile)
|
|
88
|
+
case _:
|
|
89
|
+
raise ValueError(
|
|
90
|
+
f"未知模型类型: {model_type},支持: {', '.join(_MODEL_TYPES)}"
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
@classmethod
|
|
94
|
+
def build_chat(
|
|
95
|
+
cls, model: str | None = None, *, api_key: str | None = None
|
|
96
|
+
) -> ChatService:
|
|
97
|
+
profile = cls._resolve_profile("chat", model, api_key=api_key)
|
|
98
|
+
return ChatService(profile)
|
|
99
|
+
|
|
100
|
+
@classmethod
|
|
101
|
+
def build_embedding(
|
|
102
|
+
cls, model: str | None = None, *, api_key: str | None = None
|
|
103
|
+
) -> EmbeddingService:
|
|
104
|
+
profile = cls._resolve_profile("embedding", model, api_key=api_key)
|
|
105
|
+
return EmbeddingService(profile)
|
|
106
|
+
|
|
107
|
+
@classmethod
|
|
108
|
+
def build_rerank(
|
|
109
|
+
cls, model: str | None = None, *, api_key: str | None = None
|
|
110
|
+
) -> RerankService:
|
|
111
|
+
profile = cls._resolve_profile("rerank", model, api_key=api_key)
|
|
112
|
+
return RerankService(profile)
|
|
113
|
+
|
|
114
|
+
@classmethod
|
|
115
|
+
def build_image(
|
|
116
|
+
cls, model: str | None = None, *, api_key: str | None = None
|
|
117
|
+
) -> ImageService:
|
|
118
|
+
profile = cls._resolve_profile("image", model, api_key=api_key)
|
|
119
|
+
return ImageService(profile)
|
kit/ai/model_client.py
ADDED
|
@@ -0,0 +1,214 @@
|
|
|
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
|
+
import asyncio
|
|
21
|
+
import json
|
|
22
|
+
import random
|
|
23
|
+
from collections.abc import AsyncIterator
|
|
24
|
+
from typing import Final
|
|
25
|
+
|
|
26
|
+
import httpx
|
|
27
|
+
|
|
28
|
+
from kit.ai.exceptions import (
|
|
29
|
+
AIRequestError,
|
|
30
|
+
AIResponseError,
|
|
31
|
+
AITimeoutError,
|
|
32
|
+
)
|
|
33
|
+
from kit.ai.profiles import ModelProfile
|
|
34
|
+
from kit.ai.types import JSONObject
|
|
35
|
+
|
|
36
|
+
__all__ = ["ModelClient"]
|
|
37
|
+
|
|
38
|
+
# post_json 对超时/5xx 的重试次数(指数退避)
|
|
39
|
+
_DEFAULT_RETRIES = 2
|
|
40
|
+
|
|
41
|
+
# list_models 轻探测独立超时(秒):仅验证密钥有效性,与业务请求超时无关
|
|
42
|
+
_LIST_MODELS_TIMEOUT: Final[float] = 10.0
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def _retryable(exc: BaseException) -> bool:
|
|
46
|
+
"""判断异常是否可重试:超时,或服务端 5xx(按 AI 异常族状态码判定)。"""
|
|
47
|
+
if isinstance(exc, AITimeoutError):
|
|
48
|
+
return True
|
|
49
|
+
return (
|
|
50
|
+
isinstance(exc, AIRequestError)
|
|
51
|
+
and exc.status_code is not None
|
|
52
|
+
and exc.status_code >= 500
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
class ModelClient:
|
|
57
|
+
"""模型 HTTP 客户端(全异步,单客户端)。
|
|
58
|
+
|
|
59
|
+
普通请求与流式请求共用同一个 ``httpx.AsyncClient``;流式请求在
|
|
60
|
+
调用处通过 ``timeout=None`` 局部覆盖超时,避免双客户端维护成本。
|
|
61
|
+
``post_json`` 对超时/5xx 做指数退避重试;``stream_json`` 不重试
|
|
62
|
+
(半开流连接重放语义复杂)。
|
|
63
|
+
"""
|
|
64
|
+
|
|
65
|
+
def __init__(self, profile: ModelProfile) -> None:
|
|
66
|
+
self._profile: ModelProfile = profile
|
|
67
|
+
self._client: httpx.AsyncClient = httpx.AsyncClient(
|
|
68
|
+
base_url=profile.base_url,
|
|
69
|
+
timeout=profile.timeout,
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
@property
|
|
73
|
+
def path(self) -> str:
|
|
74
|
+
return self._profile.endpoint
|
|
75
|
+
|
|
76
|
+
@property
|
|
77
|
+
def model(self) -> str:
|
|
78
|
+
return self._profile.model
|
|
79
|
+
|
|
80
|
+
@staticmethod
|
|
81
|
+
def _headers(profile: ModelProfile) -> dict[str, str]:
|
|
82
|
+
return {
|
|
83
|
+
"Content-Type": "application/json",
|
|
84
|
+
"Authorization": f"Bearer {profile.api_key}",
|
|
85
|
+
}
|
|
86
|
+
|
|
87
|
+
@staticmethod
|
|
88
|
+
async def _parse_json(data: str) -> JSONObject:
|
|
89
|
+
try:
|
|
90
|
+
parsed = json.loads(data)
|
|
91
|
+
except (json.JSONDecodeError, TypeError) as exc:
|
|
92
|
+
raise AIResponseError(f"响应不是合法 JSON: {exc}") from exc
|
|
93
|
+
if not isinstance(parsed, dict):
|
|
94
|
+
raise AIResponseError(f"响应顶层必须是 JSON 对象: {data[:200]}")
|
|
95
|
+
return parsed
|
|
96
|
+
|
|
97
|
+
@staticmethod
|
|
98
|
+
async def _parse_sse_line(line: str) -> JSONObject | None:
|
|
99
|
+
line = line.strip()
|
|
100
|
+
if not line or not line.startswith("data:"):
|
|
101
|
+
return None
|
|
102
|
+
payload = line[len("data:") :].strip()
|
|
103
|
+
if payload == "[DONE]":
|
|
104
|
+
return None
|
|
105
|
+
if payload == "[ERROR]":
|
|
106
|
+
raise AIResponseError("流式响应返回 [ERROR]")
|
|
107
|
+
return await ModelClient._parse_json(payload)
|
|
108
|
+
|
|
109
|
+
async def _post_once(self, payload: JSONObject) -> JSONObject:
|
|
110
|
+
"""执行单次请求并解析响应体(不重试)。"""
|
|
111
|
+
try:
|
|
112
|
+
resp = await self._client.post(
|
|
113
|
+
self.path,
|
|
114
|
+
headers=self._headers(self._profile),
|
|
115
|
+
json=payload,
|
|
116
|
+
)
|
|
117
|
+
resp.raise_for_status()
|
|
118
|
+
return await self._parse_json(resp.text)
|
|
119
|
+
except httpx.TimeoutException as exc:
|
|
120
|
+
raise AITimeoutError(f"请求超时: {exc}") from exc
|
|
121
|
+
except httpx.HTTPStatusError as exc:
|
|
122
|
+
raise AIRequestError(
|
|
123
|
+
f"HTTP 错误 {exc.response.status_code}: {exc.response.text[:200]}",
|
|
124
|
+
status_code=exc.response.status_code,
|
|
125
|
+
) from exc
|
|
126
|
+
except httpx.HTTPError as exc:
|
|
127
|
+
raise AIRequestError(f"请求失败: {exc}") from exc
|
|
128
|
+
|
|
129
|
+
async def post_json(
|
|
130
|
+
self,
|
|
131
|
+
payload: JSONObject,
|
|
132
|
+
*,
|
|
133
|
+
retries: int = _DEFAULT_RETRIES,
|
|
134
|
+
) -> JSONObject:
|
|
135
|
+
"""发送普通 JSON 请求并解析响应体。
|
|
136
|
+
|
|
137
|
+
对超时/5xx 做指数退避重试,累计 ``retries`` 次后抛出最后一次异常。
|
|
138
|
+
"""
|
|
139
|
+
for attempt in range(retries + 1):
|
|
140
|
+
try:
|
|
141
|
+
return await self._post_once(payload)
|
|
142
|
+
except (AITimeoutError, AIRequestError) as exc:
|
|
143
|
+
if attempt >= retries or not _retryable(exc):
|
|
144
|
+
raise
|
|
145
|
+
delay = min(2**attempt, 8) + random.uniform(0, 0.5)
|
|
146
|
+
await asyncio.sleep(delay)
|
|
147
|
+
# 不可达,供类型系统收口
|
|
148
|
+
raise AIRequestError("重试后仍未成功")
|
|
149
|
+
|
|
150
|
+
async def stream_json(
|
|
151
|
+
self, payload: JSONObject
|
|
152
|
+
) -> AsyncIterator[JSONObject | None]:
|
|
153
|
+
"""发送流式请求,逐帧产出已解析的 SSE 数据体。
|
|
154
|
+
|
|
155
|
+
空行与 ``[DONE]`` 标记会被跳过(产出 None 或不产出);
|
|
156
|
+
``[ERROR]`` 标记会抛出 ``AIResponseError``。不做重试。
|
|
157
|
+
"""
|
|
158
|
+
try:
|
|
159
|
+
async with self._client.stream(
|
|
160
|
+
"POST",
|
|
161
|
+
self.path,
|
|
162
|
+
headers=self._headers(self._profile),
|
|
163
|
+
json=payload,
|
|
164
|
+
timeout=None,
|
|
165
|
+
) as resp:
|
|
166
|
+
resp.raise_for_status()
|
|
167
|
+
async for line in resp.aiter_lines():
|
|
168
|
+
parsed = await self._parse_sse_line(line)
|
|
169
|
+
if parsed is not None:
|
|
170
|
+
yield parsed
|
|
171
|
+
except httpx.TimeoutException as exc:
|
|
172
|
+
raise AITimeoutError(f"流式请求超时: {exc}") from exc
|
|
173
|
+
except httpx.HTTPStatusError as exc:
|
|
174
|
+
raise AIRequestError(
|
|
175
|
+
f"流式 HTTP 错误 {exc.response.status_code}: {exc.response.text[:200]}",
|
|
176
|
+
status_code=exc.response.status_code,
|
|
177
|
+
) from exc
|
|
178
|
+
except httpx.HTTPError as exc:
|
|
179
|
+
raise AIRequestError(f"流式请求失败: {exc}") from exc
|
|
180
|
+
|
|
181
|
+
async def list_models(self) -> JSONObject:
|
|
182
|
+
"""GET ``{base_url}/models`` 轻探测:仅验证密钥有效性(候选链预热用)。
|
|
183
|
+
|
|
184
|
+
独立 10 秒超时(不影响业务请求配置);使用绝对 URL 直连——
|
|
185
|
+
``base_url`` 已含版本前缀,若用相对路径会被 httpx 按根路径拼接。
|
|
186
|
+
|
|
187
|
+
Raises:
|
|
188
|
+
AIRequestError: 非 2xx(``status_code`` 携带 HTTP 状态码)
|
|
189
|
+
或连接失败。
|
|
190
|
+
AITimeoutError: 探测超时。
|
|
191
|
+
AIResponseError: 响应不是合法 JSON 对象。
|
|
192
|
+
"""
|
|
193
|
+
url = self._profile.base_url.rstrip("/") + "/models"
|
|
194
|
+
try:
|
|
195
|
+
resp = await self._client.get(
|
|
196
|
+
url,
|
|
197
|
+
headers=self._headers(self._profile),
|
|
198
|
+
timeout=_LIST_MODELS_TIMEOUT,
|
|
199
|
+
)
|
|
200
|
+
resp.raise_for_status()
|
|
201
|
+
return await self._parse_json(resp.text)
|
|
202
|
+
except httpx.TimeoutException as exc:
|
|
203
|
+
raise AITimeoutError(f"模型探测超时: {exc}") from exc
|
|
204
|
+
except httpx.HTTPStatusError as exc:
|
|
205
|
+
raise AIRequestError(
|
|
206
|
+
f"模型探测 HTTP 错误 {exc.response.status_code}: "
|
|
207
|
+
f"{exc.response.text[:200]}",
|
|
208
|
+
status_code=exc.response.status_code,
|
|
209
|
+
) from exc
|
|
210
|
+
except httpx.HTTPError as exc:
|
|
211
|
+
raise AIRequestError(f"模型探测请求失败: {exc}") from exc
|
|
212
|
+
|
|
213
|
+
async def close(self) -> None:
|
|
214
|
+
await self._client.aclose()
|